From 274c3153b0aa086ec644af5de0e02208cc0f37ab Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Fri, 18 Sep 2026 16:41:20 -0700 Subject: [PATCH] Add CelLiteRuntimeVersionSkewTest for runtime version skew verification PiperOrigin-RevId: 984133626 --- .../cel/checker/CelStandardDeclarations.java | 2 + .../java/dev/cel/checker/ExprChecker.java | 9 +- .../cel/checker/CelCheckerLegacyImplTest.java | 15 + .../checker/CelStandardDeclarationsTest.java | 51 + .../java/dev/cel/checker/ExprCheckerTest.java | 4 + .../test/java/dev/cel/checker/TypesTest.java | 219 ++- .../test/resources/standardEnvDump.baseline | 6 + checker/src/test/resources/types.baseline | 23 +- .../values/RawProtoMessageLiteValue.java | 21 +- .../values/RawProtoMessageLiteValueTest.java | 228 ++++ .../java/dev/cel/optimizer/AstMutator.java | 164 +-- .../dev/cel/optimizer/AstMutatorTest.java | 45 + .../optimizers/SelectOptimizerTest.java | 8 + .../src/test/java/dev/cel/runtime/BUILD.bazel | 8 + .../CelLiteRuntimeVersionSkewTest.java | 1199 +++++++++++++++++ 15 files changed, 1894 insertions(+), 108 deletions(-) create mode 100644 runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java diff --git a/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java b/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java index 218135c1d..b02c14c5e 100644 --- a/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java +++ b/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java @@ -1553,6 +1553,8 @@ public enum StandardIdentifier { DOUBLE(newStandardIdentDecl(SimpleType.DOUBLE)), BYTES(newStandardIdentDecl(SimpleType.BYTES)), STRING(newStandardIdentDecl(SimpleType.STRING)), + DURATION(newStandardIdentDecl(SimpleType.DURATION)), + TIMESTAMP(newStandardIdentDecl(SimpleType.TIMESTAMP)), DYN(newStandardIdentDecl(SimpleType.DYN)), TYPE(newStandardIdentDecl("type", SimpleType.DYN)), NULL_TYPE(newStandardIdentDecl("null_type", SimpleType.NULL_TYPE)), diff --git a/checker/src/main/java/dev/cel/checker/ExprChecker.java b/checker/src/main/java/dev/cel/checker/ExprChecker.java index 8a842ce7a..34d0faf81 100644 --- a/checker/src/main/java/dev/cel/checker/ExprChecker.java +++ b/checker/src/main/java/dev/cel/checker/ExprChecker.java @@ -379,7 +379,8 @@ private void visit(CelMutableExpr expr, CelMutableStruct struct) { env.reportError(expr.id(), getPosition(expr), "'%s' is not a type", CelTypes.format(type)); } else { messageType = ((TypeType) type).type(); - if (!messageType.kind().equals(CelKind.STRUCT)) { + if (!messageType.kind().equals(CelKind.STRUCT) + && !CelTypes.isWellKnownType(messageType.name())) { env.reportError( expr.id(), getPosition(expr), @@ -816,7 +817,7 @@ private CelType getFieldType(long exprId, int position, CelType type, String fie // provided String errorMessage = String.format("Message type resolution failure while referencing field '%s'.", fieldName); - if (type.kind().equals(CelKind.STRUCT)) { + if (type.kind().equals(CelKind.STRUCT) || CelTypes.isWellKnownType(typeName)) { errorMessage += String.format( " Ensure that the descriptor for type '%s' was added to the environment", typeName); @@ -858,7 +859,9 @@ private static CelType normalizeFieldType(CelType celType) { /** TODO: Remove after cl/984117942 is submitted. */ private static Optional lookupLegacyFieldType( TypeProvider legacyTypeProvider, CelType type, String fieldName) { - TypeProvider.FieldType legacyFieldType = legacyTypeProvider.lookupFieldType(type, fieldName); + Type messageType = CelProtoTypes.createMessage(type.name()); + TypeProvider.FieldType legacyFieldType = + legacyTypeProvider.lookupFieldType(messageType, fieldName); if (legacyFieldType != null) { return Optional.of(legacyFieldType.celType()); } diff --git a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java index 3e557d0eb..6771c1278 100644 --- a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java +++ b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java @@ -21,6 +21,7 @@ import com.google.common.collect.ImmutableMap; import com.google.protobuf.Duration; import com.google.protobuf.FieldMask; +import com.google.protobuf.Timestamp; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.checker.CelStandardDeclarations.StandardFunction; @@ -182,6 +183,20 @@ public void check_wellKnownTypeStructCreation_withLegacyTypeProvider_success() t assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); } + @Test + public void check_wellKnownTypeTimestampStructCreation_withLegacyTypeProvider_success() + throws Exception { + TypeProvider legacyTypeProvider = + new DescriptorTypeProvider(ImmutableList.of(Timestamp.getDescriptor())); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder().setTypeProvider(legacyTypeProvider).build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Timestamp{seconds: 100, nanos: 200}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.TIMESTAMP); + } + @Test public void check_protoTypeMask_failsClosedWithLegacyTypeProvider() throws Exception { TypeProvider legacyTypeProvider = diff --git a/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java b/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java index f867728b0..867946fda 100644 --- a/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java +++ b/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java @@ -209,6 +209,18 @@ public void standardDeclarations_includeIdentifiers() { .containsExactly(StandardIdentifier.INT.identDecl(), StandardIdentifier.UINT.identDecl()); } + @Test + public void standardDeclarations_includeDurationAndTimestampIdentifiers() { + CelStandardDeclarations celStandardDeclaration = + CelStandardDeclarations.newBuilder() + .includeIdentifiers(StandardIdentifier.DURATION, StandardIdentifier.TIMESTAMP) + .build(); + + assertThat(celStandardDeclaration.identifierDecls()) + .containsExactly( + StandardIdentifier.DURATION.identDecl(), StandardIdentifier.TIMESTAMP.identDecl()); + } + @Test public void standardDeclarations_excludeIdentifiers() { CelStandardDeclarations celStandardDeclaration = @@ -222,6 +234,45 @@ public void standardDeclarations_excludeIdentifiers() { .doesNotContain(StandardIdentifier.UINT.identDecl()); } + @Test + public void standardEnvironment_excludeDurationIdentifier_compilationFails() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardDeclarations( + CelStandardDeclarations.newBuilder() + .excludeIdentifiers(StandardIdentifier.DURATION) + .build()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Duration == type(duration('1h'))").getAst()); + + assertThat(e).hasMessageThat().contains("undeclared reference to 'google'"); + } + + @Test + public void standardEnvironment_excludeTimestampIdentifier_compilationFails() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardDeclarations( + CelStandardDeclarations.newBuilder() + .excludeIdentifiers(StandardIdentifier.TIMESTAMP) + .build()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> + celCompiler + .compile("google.protobuf.Timestamp == type(timestamp('2023-01-01T00:00:00Z'))") + .getAst()); + + assertThat(e).hasMessageThat().contains("undeclared reference to 'google'"); + } + @Test public void standardDeclarations_filterIdentifiers() { CelStandardDeclarations celStandardDeclaration = diff --git a/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java b/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java index b9a7f66df..a86eaa008 100644 --- a/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java +++ b/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java @@ -812,6 +812,10 @@ public void types() throws Exception { runTest(); source = "{}.map(c,[c,type(c)])"; runTest(); + source = + "google.protobuf.Duration == type(duration('1h')) " + + "&& google.protobuf.Timestamp == type(timestamp(0))"; + runTest(); } // Enum Values diff --git a/checker/src/test/java/dev/cel/checker/TypesTest.java b/checker/src/test/java/dev/cel/checker/TypesTest.java index 786e50668..43fa36fe2 100644 --- a/checker/src/test/java/dev/cel/checker/TypesTest.java +++ b/checker/src/test/java/dev/cel/checker/TypesTest.java @@ -15,12 +15,19 @@ package dev.cel.checker; import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertThrows; import dev.cel.expr.Type; import dev.cel.expr.Type.PrimitiveType; +import com.google.protobuf.Duration; +import com.google.protobuf.Timestamp; +import com.google.testing.junit.testparameterinjector.TestParameter; +import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.common.CelAbstractSyntaxTree; +import dev.cel.common.CelContainer; import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelOverloadDecl; +import dev.cel.common.CelValidationException; import dev.cel.common.types.CelKind; import dev.cel.common.types.CelProtoTypes; import dev.cel.common.types.CelType; @@ -37,9 +44,8 @@ import java.util.Map; import org.junit.Test; import org.junit.runner.RunWith; -import org.junit.runners.JUnit4; -@RunWith(JUnit4.class) +@RunWith(TestParameterInjector.class) public class TypesTest { @Test @@ -350,6 +356,215 @@ public void compiler_typeParamInTypeType_resolvesReturnTypeString() throws Excep assertThat(ast.getResultType()).isEqualTo(SimpleType.STRING); } + private enum WellKnownTypeIdentTestCase { + DURATION_QUALIFIED("google.protobuf.Duration", TypeType.create(SimpleType.DURATION)), + DURATION_LEADING_DOT(".google.protobuf.Duration", TypeType.create(SimpleType.DURATION)), + DURATION_UNQUALIFIED("Duration", TypeType.create(SimpleType.DURATION)), + TIMESTAMP_QUALIFIED("google.protobuf.Timestamp", TypeType.create(SimpleType.TIMESTAMP)), + TIMESTAMP_LEADING_DOT(".google.protobuf.Timestamp", TypeType.create(SimpleType.TIMESTAMP)), + TIMESTAMP_UNQUALIFIED("Timestamp", TypeType.create(SimpleType.TIMESTAMP)); + + private final String expression; + private final CelType expectedType; + + WellKnownTypeIdentTestCase(String expression, CelType expectedType) { + this.expression = expression; + this.expectedType = expectedType; + } + } + + @Test + public void compiler_wellKnownProtoTypeIdent_resolvesToSimpleType( + @TestParameter WellKnownTypeIdentTestCase testCase) throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor(), Timestamp.getDescriptor()) + .setContainer(CelContainer.ofName("google.protobuf")) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile(testCase.expression).getAst(); + + assertThat(ast.getResultType()).isEqualTo(testCase.expectedType); + } + + private enum WellKnownTypeParamTestCase { + DURATION_QUALIFIED_CAST("cast('1h', google.protobuf.Duration)", SimpleType.DURATION), + DURATION_QUALIFIED_EQUALITY( + "cast('1h', google.protobuf.Duration) == duration('1h')", SimpleType.BOOL), + DURATION_UNQUALIFIED_COMPARISON("cast('1h', Duration) > duration('0s')", SimpleType.BOOL), + TIMESTAMP_QUALIFIED_CAST("cast(0, google.protobuf.Timestamp)", SimpleType.TIMESTAMP), + TIMESTAMP_QUALIFIED_EQUALITY( + "cast(0, google.protobuf.Timestamp) == timestamp(0)", SimpleType.BOOL), + TIMESTAMP_UNQUALIFIED_COMPARISON("cast(0, Timestamp) > timestamp(0)", SimpleType.BOOL); + + private final String expression; + private final CelType expectedType; + + WellKnownTypeParamTestCase(String expression, CelType expectedType) { + this.expression = expression; + this.expectedType = expectedType; + } + } + + @Test + public void compiler_typeParamInTypeType_withWellKnownProto_resolvesWellKnownOverloads( + @TestParameter WellKnownTypeParamTestCase testCase) throws Exception { + TypeParamType typeParamT = TypeParamType.create("T"); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor(), Timestamp.getDescriptor()) + .setContainer(CelContainer.ofName("google.protobuf")) + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "cast", + CelOverloadDecl.newGlobalOverload( + "cast_t", typeParamT, SimpleType.DYN, TypeType.create(typeParamT)))) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile(testCase.expression).getAst(); + + assertThat(ast.getResultType()).isEqualTo(testCase.expectedType); + } + + @Test + public void compiler_durationIdent_withoutMessageTypes_resolvesToSimpleType() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Duration").getAst(); + + assertThat(ast.getResultType()).isEqualTo(TypeType.create(SimpleType.DURATION)); + } + + @Test + public void compiler_timestampIdent_withoutMessageTypes_resolvesToSimpleType() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Timestamp").getAst(); + + assertThat(ast.getResultType()).isEqualTo(TypeType.create(SimpleType.TIMESTAMP)); + } + + @Test + public void compiler_durationStructCreation_withDescriptor_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Duration{seconds: 10, nanos: 20}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); + } + + @Test + public void compiler_timestampStructCreation_withDescriptor_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Timestamp.getDescriptor()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Timestamp{seconds: 100, nanos: 200}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.TIMESTAMP); + } + + @Test + public void compiler_durationStructCreation_emptyFields_success() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Duration{}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); + } + + @Test + public void compiler_timestampStructCreation_emptyFields_success() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Timestamp{}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.TIMESTAMP); + } + + @Test + public void compiler_durationStructCreation_withoutDescriptor_throws() { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Duration{seconds: 10}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains( + "Message type resolution failure while referencing field 'seconds'." + + " Ensure that the descriptor for type 'google.protobuf.Duration' was added to the" + + " environment"); + } + + @Test + public void compiler_timestampStructCreation_withoutDescriptor_throws() { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Timestamp{seconds: 10}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains( + "Message type resolution failure while referencing field 'seconds'. Ensure that the" + + " descriptor for type 'google.protobuf.Timestamp' was added to the environment"); + } + + @Test + public void compiler_durationStructCreation_typeMismatch_throws() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Duration{seconds: 'bad'}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains("expected type of field 'seconds' is 'int' but provided type is 'string'"); + } + + @Test + public void compiler_timestampStructCreation_typeMismatch_throws() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Timestamp.getDescriptor()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Timestamp{seconds: 'bad'}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains("expected type of field 'seconds' is 'int' but provided type is 'string'"); + } + + @Test + public void compiler_structCreation_primitiveType_throws() { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelValidationException e = + assertThrows(CelValidationException.class, () -> celCompiler.compile("int{}").getAst()); + + assertThat(e).hasMessageThat().contains("'int' is not a message type"); + } + @Test public void compiler_typeParamInCompositeTypeType_resolvesReturnType() throws Exception { TypeParamType typeParamT = TypeParamType.create("T"); diff --git a/checker/src/test/resources/standardEnvDump.baseline b/checker/src/test/resources/standardEnvDump.baseline index 864bfd340..6b8f716fb 100644 --- a/checker/src/test/resources/standardEnvDump.baseline +++ b/checker/src/test/resources/standardEnvDump.baseline @@ -228,6 +228,12 @@ declare getSeconds { function timestamp_to_seconds_with_tz google.protobuf.Timestamp.(string) -> int function duration_to_seconds google.protobuf.Duration.() -> int } +declare google.protobuf.Duration { + value type(google.protobuf.Duration) +} +declare google.protobuf.Timestamp { + value type(google.protobuf.Timestamp) +} declare int { value type(int) } diff --git a/checker/src/test/resources/types.baseline b/checker/src/test/resources/types.baseline index 939e0ed97..f88da0a24 100644 --- a/checker/src/test/resources/types.baseline +++ b/checker/src/test/resources/types.baseline @@ -45,4 +45,25 @@ __comprehension__( ]~list(list(dyn)) )~list(list(dyn))^add_list, // Result - @result~list(list(dyn))^@result)~list(list(dyn)) \ No newline at end of file + @result~list(list(dyn))^@result)~list(list(dyn)) + +Source: google.protobuf.Duration == type(duration('1h')) && google.protobuf.Timestamp == type(timestamp(0)) +=====> +_&&_( + _==_( + google.protobuf.Duration~type(google.protobuf.Duration)^google.protobuf.Duration, + type( + duration( + "1h"~string + )~google.protobuf.Duration^string_to_duration + )~type(google.protobuf.Duration)^type + )~bool^equals, + _==_( + google.protobuf.Timestamp~type(google.protobuf.Timestamp)^google.protobuf.Timestamp, + type( + timestamp( + 0~int + )~google.protobuf.Timestamp^int64_to_timestamp + )~type(google.protobuf.Timestamp)^type + )~bool^equals +)~bool^logical_and \ No newline at end of file diff --git a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java index 2a3bdf940..dc5f24072 100644 --- a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java @@ -41,6 +41,8 @@ import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor.EncodingType; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor.JavaType; import java.io.IOException; +import java.time.Duration; +import java.time.Instant; import java.util.ArrayList; import java.util.LinkedHashMap; import java.util.List; @@ -193,14 +195,25 @@ private static Object decodeWireField( fieldDescriptor != null ? fieldDescriptor.getEncodingType() == EncodingType.LIST : field.defaultValue() instanceof List; - String protoTypeName = - fieldDescriptor != null - ? fieldDescriptor.getFieldProtoTypeName() - : UNKNOWN_MESSAGE_TYPE_NAME; + String protoTypeName = resolveProtoTypeName(field, fieldDescriptor); return decodeWireEntries(unknowns, typeCode, protoTypeName, isRepeated, converter); } + private static String resolveProtoTypeName( + SelectField field, @Nullable FieldLiteDescriptor fieldDescriptor) { + if (fieldDescriptor != null) { + return fieldDescriptor.getFieldProtoTypeName(); + } + if (field.defaultValue() instanceof Duration) { + return WellKnownProto.DURATION.typeName(); + } + if (field.defaultValue() instanceof Instant) { + return WellKnownProto.TIMESTAMP.typeName(); + } + return UNKNOWN_MESSAGE_TYPE_NAME; + } + private static Object resolveDefault( SelectField field, @Nullable FieldLiteDescriptor fieldDescriptor, diff --git a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java index 9a4079a4c..9ef830998 100644 --- a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java @@ -40,6 +40,7 @@ import dev.cel.protobuf.CelLiteDescriptor.MessageLiteDescriptor; import java.io.ByteArrayOutputStream; import java.time.Duration; +import java.time.Instant; import java.util.NoSuchElementException; import java.util.Optional; import org.junit.Test; @@ -1260,4 +1261,231 @@ public void selectByFieldNumber_absentMessageFieldWithoutDescriptor_returnsUnkno .hasMessageThat() .contains("Decoding unknown map field from wire bytes is unsupported"); } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum WellKnownFieldWithoutDescriptorTestCase { + DURATION( + TestAllTypes.newBuilder() + .setSingleDuration(ProtoTimeUtils.toProtoDuration(Duration.ofSeconds(120L, 500L))) + .build(), + SelectField.create( + TestAllTypes.SINGLE_DURATION_FIELD_NUMBER, + "single_duration", + FieldLiteDescriptor.Type.MESSAGE.getNumber(), + Duration.ZERO), + Duration.ofSeconds(120L, 500L), + Duration.ZERO), + TIMESTAMP( + TestAllTypes.newBuilder() + .setSingleTimestamp( + ProtoTimeUtils.toProtoTimestamp(Instant.ofEpochSecond(1700000000L, 123456789L))) + .build(), + SelectField.create( + TestAllTypes.SINGLE_TIMESTAMP_FIELD_NUMBER, + "single_timestamp", + FieldLiteDescriptor.Type.MESSAGE.getNumber(), + Instant.EPOCH), + Instant.ofEpochSecond(1700000000L, 123456789L), + Instant.EPOCH); + + private final TestAllTypes populatedProto; + private final SelectField selectField; + private final Object expectedPopulatedValue; + private final Object expectedDefaultValue; + + WellKnownFieldWithoutDescriptorTestCase( + TestAllTypes populatedProto, + SelectField selectField, + Object expectedPopulatedValue, + Object expectedDefaultValue) { + this.populatedProto = populatedProto; + this.selectField = selectField; + this.expectedPopulatedValue = expectedPopulatedValue; + this.expectedDefaultValue = expectedDefaultValue; + } + } + + @Test + public void selectByFieldNumber_wellKnownFieldWithoutDescriptor_decodesOrReturnsDefault( + @TestParameter WellKnownFieldWithoutDescriptorTestCase testCase) { + RawProtoMessageLiteValue populatedRaw = + RawProtoMessageLiteValue.create( + testCase.populatedProto.toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + RawProtoMessageLiteValue emptyRaw = + RawProtoMessageLiteValue.create( + TestAllTypes.getDefaultInstance().toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + + Object populatedSelected = populatedRaw.selectByFieldNumber(testCase.selectField); + Object emptySelected = emptyRaw.selectByFieldNumber(testCase.selectField); + + assertThat(populatedSelected).isEqualTo(testCase.expectedPopulatedValue); + assertThat(emptySelected).isEqualTo(testCase.expectedDefaultValue); + } + + @Test + public void findByFieldNumber_wellKnownFieldWithoutDescriptor_returnsOptionalValue( + @TestParameter WellKnownFieldWithoutDescriptorTestCase testCase) { + RawProtoMessageLiteValue populatedRaw = + RawProtoMessageLiteValue.create( + testCase.populatedProto.toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + RawProtoMessageLiteValue emptyRaw = + RawProtoMessageLiteValue.create( + TestAllTypes.getDefaultInstance().toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + + Optional populatedFound = populatedRaw.findByFieldNumber(testCase.selectField); + Optional emptyFound = emptyRaw.findByFieldNumber(testCase.selectField); + + assertThat(populatedFound).hasValue(testCase.expectedPopulatedValue); + assertThat(emptyFound).isEmpty(); + } + + @Test + public void selectByFieldNumber_negativeInt32Varint_decodesSignedIntCorrectly() { + TestAllTypes proto = TestAllTypes.newBuilder().setSingleInt32(-42).build(); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + proto.toByteString(), "cel.expr.conformance.proto3.TestAllTypes", EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.SINGLE_INT32_FIELD_NUMBER, + "single_int32", + FieldLiteDescriptor.Type.INT32.getNumber(), + 0L); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isEqualTo(-42L); + } + + @Test + public void selectByFieldNumber_negativeInt32MinValue_decodesSignedIntCorrectly() { + TestAllTypes proto = TestAllTypes.newBuilder().setSingleInt32(Integer.MIN_VALUE).build(); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + proto.toByteString(), "cel.expr.conformance.proto3.TestAllTypes", EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.SINGLE_INT32_FIELD_NUMBER, + "single_int32", + FieldLiteDescriptor.Type.INT32.getNumber(), + 0L); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isEqualTo((long) Integer.MIN_VALUE); + } + + @Test + public void selectByFieldNumber_unpackedAndPackedRepeatedInt64_decodeIdentically() + throws Exception { + ByteArrayOutputStream unpackedBaos = new ByteArrayOutputStream(); + CodedOutputStream unpackedCos = CodedOutputStream.newInstance(unpackedBaos); + unpackedCos.writeInt64(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, 10L); + unpackedCos.writeInt64(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, 20L); + unpackedCos.writeInt64(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, 30L); + unpackedCos.flush(); + RawProtoMessageLiteValue unpackedRaw = + RawProtoMessageLiteValue.create( + ByteString.copyFrom(unpackedBaos.toByteArray()), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + TestAllTypes packedProto = + TestAllTypes.newBuilder() + .addRepeatedInt64(10L) + .addRepeatedInt64(20L) + .addRepeatedInt64(30L) + .build(); + RawProtoMessageLiteValue packedRaw = + RawProtoMessageLiteValue.create( + packedProto.toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.REPEATED_INT64_FIELD_NUMBER, + "repeated_int64", + FieldLiteDescriptor.Type.INT64.getNumber(), + ImmutableList.of()); + + Object unpackedResult = unpackedRaw.selectByFieldNumber(field); + Object packedResult = packedRaw.selectByFieldNumber(field); + + assertThat(unpackedResult).isEqualTo(ImmutableList.of(10L, 20L, 30L)); + assertThat(packedResult).isEqualTo(ImmutableList.of(10L, 20L, 30L)); + assertThat( + unpackedRaw.hasFieldByNumber( + SelectField.create(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, "repeated_int64"))) + .isTrue(); + assertThat( + packedRaw.hasFieldByNumber( + SelectField.create(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, "repeated_int64"))) + .isTrue(); + assertThat(unpackedRaw.findByFieldNumber(field)).hasValue(ImmutableList.of(10L, 20L, 30L)); + assertThat(packedRaw.findByFieldNumber(field)).hasValue(ImmutableList.of(10L, 20L, 30L)); + } + + @Test + public void selectByFieldNumber_mixedUnpackedAndPackedRepeatedInt64_concatenatesInOrder() + throws Exception { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + CodedOutputStream cos = CodedOutputStream.newInstance(baos); + cos.writeInt64(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, 10L); + ByteArrayOutputStream packedChunk = new ByteArrayOutputStream(); + CodedOutputStream packedCos = CodedOutputStream.newInstance(packedChunk); + packedCos.writeInt64NoTag(20L); + packedCos.writeInt64NoTag(30L); + packedCos.flush(); + cos.writeByteArray(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, packedChunk.toByteArray()); + cos.writeInt64(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, 40L); + cos.flush(); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + ByteString.copyFrom(baos.toByteArray()), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.REPEATED_INT64_FIELD_NUMBER, + "repeated_int64", + FieldLiteDescriptor.Type.INT64.getNumber(), + ImmutableList.of()); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isEqualTo(ImmutableList.of(10L, 20L, 30L, 40L)); + } + + @Test + public void selectByFieldNumber_fiveByteUnsignedInt32Varint_signExtendsCorrectly() { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + baos.write((TestAllTypes.SINGLE_INT32_FIELD_NUMBER << 3)); + baos.write(0xD6); + baos.write(0xFF); + baos.write(0xFF); + baos.write(0xFF); + baos.write(0x0F); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + ByteString.copyFrom(baos.toByteArray()), + "cel.expr.conformance.proto3.TestAllTypes", + EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.SINGLE_INT32_FIELD_NUMBER, + "single_int32", + FieldLiteDescriptor.Type.INT32.getNumber(), + 0L); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isEqualTo(-42L); + } } diff --git a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java index c770844ff..2e3027dcc 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java +++ b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java @@ -14,6 +14,7 @@ package dev.cel.optimizer; +import static com.google.common.collect.ImmutableList.toImmutableList; import static com.google.common.collect.ImmutableMap.toImmutableMap; import static java.lang.Math.max; import static java.util.stream.Collectors.toCollection; @@ -23,6 +24,7 @@ import com.google.common.base.Preconditions; import com.google.common.base.Strings; import com.google.common.collect.HashBasedTable; +import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; import com.google.common.collect.Streams; import com.google.common.collect.Table; @@ -46,14 +48,12 @@ import java.util.Arrays; import java.util.Collection; import java.util.HashMap; -import java.util.LinkedHashMap; import java.util.List; import java.util.Map.Entry; import java.util.NoSuchElementException; import java.util.Optional; import java.util.function.Function; import java.util.function.Predicate; -import java.util.stream.Collectors; /** AstMutator contains logic for mutating a {@link CelAbstractSyntaxTree}. */ @Immutable @@ -187,8 +187,6 @@ public CelMutableAst renumberIdsConsecutively(CelMutableAst mutableAst) { * *

The expression IDs are not modified when the identifier names are changed. * - *

Mangling occurs only if the iteration variable is referenced within the loop step. - * *

Iteration variables in comprehensions are numbered based on their comprehension nesting * levels and the iteration variable's type. Examples: * @@ -222,98 +220,14 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames( .and( node -> !node.expr().comprehension().iterVar2().startsWith(newIterVar2Prefix + ":")); - LinkedHashMap comprehensionsToMangle = + ImmutableList comprehensionsToMangle = navigableMutableAst .getRoot() // This is important - mangling needs to happen bottom-up to avoid stepping over // shadowed variables that are not part of the comprehension being mangled. .allNodes(TraversalOrder.POST_ORDER) .filter(comprehensionIdentifierPredicate) - .filter( - node -> { - // Ensure the iter_var or the comprehension result is actually referenced in the - // loop_step. If it's not, we can skip mangling. - String iterVar = node.expr().comprehension().iterVar(); - String iterVar2 = node.expr().comprehension().iterVar2(); - String result = node.expr().comprehension().result().ident().name(); - return CelNavigableMutableExpr.fromExpr(node.expr().comprehension().loopStep()) - .allNodes() - .filter(subNode -> subNode.getKind().equals(ExprKind.Kind.IDENT)) - .map(subNode -> subNode.expr().ident()) - .anyMatch( - ident -> - ident.name().contains(iterVar) - || ident.name().contains(iterVar2) - || ident.name().contains(result)); - }) - .collect( - Collectors.toMap( - k -> k, - v -> { - CelMutableComprehension comprehension = v.expr().comprehension(); - String iterVar = comprehension.iterVar(); - String iterVar2 = comprehension.iterVar2(); - // Identifiers to mangle could be the iteration variable, comprehension - // result or both, but at least one has to exist. - // As an example, [1,2].map(i, 3) would result in optional.empty for iteration - // variable because `i` is not actually used. - Optional iterVarId = - CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) - .allNodes() - .filter( - loopStepNode -> - loopStepNode.getKind().equals(ExprKind.Kind.IDENT) - && loopStepNode.expr().ident().name().equals(iterVar)) - .map(CelNavigableMutableExpr::id) - .findAny(); - Optional iterVar2Id = - CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) - .allNodes() - .filter( - loopStepNode -> - !iterVar2.isEmpty() - && loopStepNode.getKind().equals(ExprKind.Kind.IDENT) - && loopStepNode.expr().ident().name().equals(iterVar2)) - .map(CelNavigableMutableExpr::id) - .findAny(); - Optional iterVarType = - iterVarId.map( - id -> - navigableMutableAst - .getType(id) - .orElseThrow( - () -> - new NoSuchElementException( - "Checked type not present for iteration" - + " variable: " - + iterVarId))); - Optional iterVar2Type = - iterVar2Id.map( - id -> - navigableMutableAst - .getType(id) - .orElseThrow( - () -> - new NoSuchElementException( - "Checked type not present for iteration" - + " variable: " - + iterVar2Id))); - CelType resultType = - navigableMutableAst - .getType(comprehension.result().id()) - .orElseThrow( - () -> - new IllegalStateException( - "Result type was not present for the comprehension ID: " - + comprehension.result().id())); - - return MangledComprehensionType.of(iterVarType, iterVar2Type, resultType); - }, - (x, y) -> { - throw new IllegalStateException( - "Unexpected CelNavigableMutableExpr collision"); - }, - LinkedHashMap::new)); + .collect(toImmutableList()); // The map that we'll eventually return to the caller. HashMap mangledIdentNamesToType = @@ -324,12 +238,10 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames( CelMutableExpr mutatedComprehensionExpr = navigableMutableAst.getAst().expr(); CelMutableSource newSource = navigableMutableAst.getAst().source(); int iterCount = 0; - for (Entry comprehensionEntry : - comprehensionsToMangle.entrySet()) { - CelNavigableMutableExpr comprehensionNode = comprehensionEntry.getKey(); - MangledComprehensionType comprehensionEntryType = comprehensionEntry.getValue(); - + for (CelNavigableMutableExpr comprehensionNode : comprehensionsToMangle) { CelMutableExpr comprehensionExpr = comprehensionNode.expr(); + MangledComprehensionType comprehensionEntryType = + resolveComprehensionType(navigableMutableAst, comprehensionExpr.comprehension()); MangledComprehensionName mangledComprehensionName = getMangledComprehensionName( newIterVarPrefix, @@ -375,6 +287,63 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames( ImmutableMap.copyOf(mangledIdentNamesToType)); } + private static MangledComprehensionType resolveComprehensionType( + CelNavigableMutableAst navigableMutableAst, CelMutableComprehension comprehension) { + String iterVar = comprehension.iterVar(); + String iterVar2 = comprehension.iterVar2(); + // Identifiers to mangle could be the iteration variable, comprehension + // result or both, but at least one has to exist. + // As an example, [1,2].map(i, 3) would result in optional.empty for iteration + // variable because `i` is not actually used. + Optional iterVarId = + CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) + .allNodes() + .filter( + loopStepNode -> + loopStepNode.getKind().equals(ExprKind.Kind.IDENT) + && loopStepNode.expr().ident().name().equals(iterVar)) + .map(CelNavigableMutableExpr::id) + .findAny(); + Optional iterVar2Id = + CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) + .allNodes() + .filter( + loopStepNode -> + !iterVar2.isEmpty() + && loopStepNode.getKind().equals(ExprKind.Kind.IDENT) + && loopStepNode.expr().ident().name().equals(iterVar2)) + .map(CelNavigableMutableExpr::id) + .findAny(); + Optional iterVarType = + iterVarId.map( + id -> + navigableMutableAst + .getType(id) + .orElseThrow( + () -> + new NoSuchElementException( + "Checked type not present for iteration variable: " + id))); + Optional iterVar2Type = + iterVar2Id.map( + id -> + navigableMutableAst + .getType(id) + .orElseThrow( + () -> + new NoSuchElementException( + "Checked type not present for iteration variable: " + id))); + CelType resultType = + navigableMutableAst + .getType(comprehension.accuInit().id()) + .orElseThrow( + () -> + new IllegalStateException( + "Result type was not present for the comprehension ID: " + + comprehension.accuInit().id())); + + return MangledComprehensionType.of(iterVarType, iterVar2Type, resultType); + } + private static MangledComprehensionName getMangledComprehensionName( String newIterVarPrefix, String newIterVar2Prefix, @@ -1046,9 +1015,8 @@ private static long getMaxId(CelMutableAst mutableAst) { private static long getMaxId(CelNavigableMutableAst navAst) { long maxId = navAst.getRoot().maxId(); - for (Entry macroCall : - navAst.getAst().source().getMacroCalls().entrySet()) { - maxId = max(maxId, getMaxId(macroCall.getValue())); + for (CelMutableExpr macroCall : navAst.getAst().source().getMacroCalls().values()) { + maxId = max(maxId, getMaxId(macroCall)); } return maxId; diff --git a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java index f0c3a7045..2138ff142 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java @@ -868,6 +868,27 @@ public void mangleComprehensionVariable_adjacentMacros_differentIterVarTypes() t assertConsistentMacroCalls(ast); } + @Test + public void mangleComprehensionVariable_nonIdentComprehensionResult() throws Exception { + // exists_one expands into a comprehension whose result is `@result == 1` (bool) rather than a + // bare accumulator identifier, while its accumulator is initialized to `0` (int). + CelAbstractSyntaxTree ast = + CEL.compile("[1, 2, 3].exists_one(i, i > 2) && [1, 2, 3].exists(j, j > 2)").getAst(); + + CelAbstractSyntaxTree mangledAst = + AST_MUTATOR + .mangleComprehensionIdentifierNames(CelMutableAst.fromCelAst(ast), "@it", "@it2", "@ac") + .mutableAst() + .toParsedAst(); + + assertThat(CEL_UNPARSER.unparse(mangledAst)) + .isEqualTo( + "[1, 2, 3].exists_one(@it:0:0, @it:0:0 > 2) && [1, 2, 3].exists(@it:0:1, @it:0:1 >" + + " 2)"); + assertThat(CEL.createProgram(CEL.check(mangledAst).getAst()).eval()).isEqualTo(true); + assertConsistentMacroCalls(mangledAst); + } + @Test public void mangleComprehensionVariable_macroSourceDisabled_macroCallMapIsEmpty() throws Exception { @@ -1011,6 +1032,30 @@ public void mangleComprehensionVariable_nestedMacroWithShadowedVariables() throw assertConsistentMacroCalls(ast); } + @Test + public void mangleComprehensionVariable_nestedMacroWithShadowedVariables_differentTypes() + throws Exception { + CelAbstractSyntaxTree ast = + CEL.compile( + "['a', 'b'].exists(x, [1, 2].exists(x, x > 0) && x == 'a') && " + + "[1, 2].exists(x, [1, 2].exists(x, x > 0) && x == 1)") + .getAst(); + + CelAbstractSyntaxTree mangledAst = + AST_MUTATOR + .mangleComprehensionIdentifierNames(CelMutableAst.fromCelAst(ast), "@it", "@it2", "@ac") + .mutableAst() + .toParsedAst(); + + assertThat(CEL_UNPARSER.unparse(mangledAst)) + .isEqualTo( + "[\"a\", \"b\"].exists(@it:1:0, [1, 2].exists(@it:0:0, @it:0:0 > 0) && @it:1:0 ==" + + " \"a\") && [1, 2].exists(@it:1:1, [1, 2].exists(@it:0:0, @it:0:0 > 0) &&" + + " @it:1:1 == 1)"); + assertThat(CEL.createProgram(CEL.check(mangledAst).getAst()).eval()).isEqualTo(true); + assertConsistentMacroCalls(mangledAst); + } + @Test public void mangleComprehensionVariable_hasMacro_noOp() throws Exception { CelAbstractSyntaxTree ast = CEL.compile("has(msg.single_int64)").getAst(); diff --git a/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java b/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java index 3b503f301..b1cc01de4 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java @@ -300,6 +300,14 @@ private enum RewriteTestCase { "msg.single_duration", "cel.@attribute(msg, [[101, \"single_duration\", 11, duration(\"0s\")]]," + " google.protobuf.Duration)"), + PROTO3_TIMESTAMP_COMPARISON( + "msg.single_timestamp > timestamp(0)", + "cel.@attribute(msg, [[102, \"single_timestamp\", 11, timestamp(0)]]," + + " google.protobuf.Timestamp) > timestamp(0)"), + PROTO3_DURATION_COMPARISON( + "msg.single_duration == duration(\"1h\")", + "cel.@attribute(msg, [[101, \"single_duration\", 11, duration(\"0s\")]]," + + " google.protobuf.Duration) == duration(\"1h\")"), // Map selects MAP_FIELD_INDEXING( diff --git a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel index 4ea324c96..790732629 100644 --- a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel +++ b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel @@ -40,6 +40,7 @@ java_library( "//common:options", "//common:proto_v1alpha1_ast", "//common/ast", + "//common/ast:cel_block", "//common/exceptions:bad_format", "//common/exceptions:divide_by_zero", "//common/exceptions:numeric_overflow", @@ -57,13 +58,19 @@ java_library( "//common/types:message_type_provider", "//common/values", "//common/values:cel_byte_string", + "//common/values:proto_message_lite_value", "//common/values:proto_message_lite_value_provider", "//compiler", "//compiler:compiler_builder", "//extensions", "//extensions:optional_library", + "//optimizer", + "//optimizer:optimizer_builder", + "//optimizer/optimizers:common_subexpression_elimination", + "//optimizer/optimizers:select_optimizer", "//parser:macro", "//parser:unparser", + "//protobuf:cel_lite_descriptor", "//runtime", "//runtime:accumulated_unknowns", "//runtime:activation", @@ -79,6 +86,7 @@ java_library( "//runtime:lite_runtime", "//runtime:lite_runtime_factory", "//runtime:partial_vars", + "//runtime:program", "//runtime:proto_message_activation_factory", "//runtime:proto_message_runtime_equality", "//runtime:proto_message_runtime_helpers", diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java new file mode 100644 index 000000000..176f75b3b --- /dev/null +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java @@ -0,0 +1,1199 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dev.cel.runtime; + +import static com.google.common.truth.Truth.assertThat; +import static java.nio.charset.StandardCharsets.UTF_8; +import static org.junit.Assert.assertThrows; + +import com.google.api.expr.v1alpha1.CheckedExpr; +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; +import com.google.common.primitives.UnsignedLong; +import com.google.protobuf.ByteString; +import com.google.protobuf.CodedOutputStream; +import com.google.protobuf.ExtensionRegistryLite; +import com.google.testing.junit.testparameterinjector.TestParameter; +import com.google.testing.junit.testparameterinjector.TestParameterInjector; +import dev.cel.bundle.Cel; +import dev.cel.bundle.CelFactory; +import dev.cel.common.CelAbstractSyntaxTree; +import dev.cel.common.CelContainer; +import dev.cel.common.CelOptions; +import dev.cel.common.CelProtoV1Alpha1AbstractSyntaxTree; +import dev.cel.common.ast.CelBlock; +import dev.cel.common.internal.ProtoTimeUtils; +import dev.cel.common.types.StructTypeReference; +import dev.cel.common.values.CelByteString; +import dev.cel.common.values.ProtoMessageLiteValueProvider; +import dev.cel.common.values.RawProtoMessageLiteValue; +import dev.cel.expr.conformance.proto3.NestedTestAllTypes; +import dev.cel.expr.conformance.proto3.NestedTestAllTypesCelDescriptor; +import dev.cel.expr.conformance.proto3.TestAllTypes; +import dev.cel.expr.conformance.proto3.TestAllTypes.NestedEnum; +import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage; +import dev.cel.expr.conformance.proto3.TestAllTypesCelDescriptor; +import dev.cel.extensions.CelExtensions; +import dev.cel.optimizer.CelOptimizer; +import dev.cel.optimizer.CelOptimizerFactory; +import dev.cel.optimizer.optimizers.SelectOptimizer; +import dev.cel.optimizer.optimizers.SelectOptimizer.SelectOptimizerOptions; +import dev.cel.optimizer.optimizers.SubexpressionOptimizer; +import dev.cel.optimizer.optimizers.SubexpressionOptimizer.SubexpressionOptimizerOptions; +import dev.cel.parser.CelStandardMacro; +import dev.cel.protobuf.CelLiteDescriptor; +import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; +import dev.cel.protobuf.CelLiteDescriptor.MessageLiteDescriptor; +import java.io.ByteArrayOutputStream; +import java.time.Duration; +import java.time.Instant; +import java.util.List; +import org.jspecify.annotations.Nullable; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; + +@RunWith(TestParameterInjector.class) +public final class CelLiteRuntimeVersionSkewTest { + + private static final CelContainer CEL_CONTAINER = + CelContainer.ofName("cel.expr.conformance.proto3"); + + private static final CelOptions CEL_OPTIONS = + CelOptions.current() + .populateMacroCalls(true) + .enableHeterogeneousNumericComparisons(true) + .build(); + + private static final ImmutableSet SKEW_EXCLUDED_FIELD_NAMES = + ImmutableSet.of( + "single_int64", + "single_uint32", + "single_uint64", + "single_sint32", + "single_sint64", + "single_fixed32", + "single_fixed64", + "single_sfixed32", + "single_sfixed64", + "single_float", + "single_double", + "single_bool", + "single_string", + "single_bytes", + "single_nested_message", + "single_nested_enum", + "standalone_enum", + "single_duration", + "single_timestamp", + "oneof_type", + "oneof_bool", + "repeated_int64", + "repeated_double", + "repeated_string", + "repeated_nested_message", + "repeated_nested_enum", + "map_int32_int32"); + + private static final ImmutableMap SKEW_RENAMED_FIELD_NAMES = + ImmutableMap.of( + "single_int32", "v1_single_int32", + "standalone_message", "v1_standalone_message", + "repeated_int32", "v1_repeated_int32", + "map_string_string", "v1_map_string_string", + "map_int64_message", "v1_map_int64_message"); + + private static final TestAllTypes POPULATED_V2_MESSAGE = + TestAllTypes.newBuilder() + .setSingleInt64(-42L) + .setSingleUint32(123) + .setSingleUint64(999L) + .setSingleSint32(-15) + .setSingleSint64(-250L) + .setSingleFixed32(320) + .setSingleFixed64(640L) + .setSingleSfixed32(-32) + .setSingleSfixed64(-64L) + .setSingleFloat(1.5f) + .setSingleDouble(0.85d) + .setSingleBool(true) + .setSingleString("cel-skew-test") + .setSingleBytes(ByteString.copyFromUtf8("binary")) + .setSingleNestedMessage(NestedMessage.newBuilder().setBb(123).build()) + .setStandaloneEnum(NestedEnum.BAZ) + .setSingleDuration(ProtoTimeUtils.toProtoDuration(Duration.ofHours(1))) + .setSingleTimestamp( + ProtoTimeUtils.toProtoTimestamp(Instant.ofEpochSecond(1700000000L, 500L))) + .setOneofBool(true) + .addRepeatedInt64(10L) + .addRepeatedInt64(20L) + .addRepeatedDouble(0.25d) + .addRepeatedDouble(0.75d) + .addRepeatedString("foo") + .addRepeatedString("bar") + .addRepeatedNestedEnum(NestedEnum.BAR) + .addRepeatedNestedEnum(NestedEnum.BAZ) + .addRepeatedNestedMessage(NestedMessage.newBuilder().setBb(10).build()) + .addRepeatedNestedMessage(NestedMessage.newBuilder().setBb(20).build()) + .putMapInt32Int32(1, 2) + .build(); + + private static final TestAllTypes POPULATED_RENAMED_MESSAGE = + TestAllTypes.newBuilder() + .setSingleInt32(42) + .setStandaloneMessage(NestedMessage.newBuilder().setBb(77).build()) + .addRepeatedInt32(10) + .addRepeatedInt32(20) + .putMapStringString("k", "v") + .putMapInt64Message(1L, NestedMessage.newBuilder().setBb(100).build()) + .build(); + + private Cel cel; + private CelOptimizer celOptimizer; + private CelLiteRuntime v1Runtime; + + @Before + public void setUp() { + // Schema V2 Compiler: compiles policies against full V2 descriptors. + cel = + CelFactory.standardCelBuilder() + .setOptions(CEL_OPTIONS) + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .addCompilerLibraries(CelExtensions.bindings()) + .addMessageTypes(TestAllTypes.getDescriptor(), NestedTestAllTypes.getDescriptor()) + .addVar("msg", StructTypeReference.create(TestAllTypes.getDescriptor().getFullName())) + .addVar( + "nested_msg", + StructTypeReference.create(NestedTestAllTypes.getDescriptor().getFullName())) + .setContainer(CEL_CONTAINER) + .build(); + + celOptimizer = + CelOptimizerFactory.standardCelOptimizerBuilder(cel) + .addAstOptimizers( + SelectOptimizer.newInstance( + SelectOptimizerOptions.newBuilder().build(), + TestAllTypes.getDescriptor().getFile())) + .build(); + + // Schema V1 Runtime: restricted CelLiteDescriptor omitting fields in SKEW_EXCLUDED_FIELD_NAMES + // and renaming fields in SKEW_RENAMED_FIELD_NAMES while preserving their field numbers. + // Configured without ProtoMessageTypeProvider to match Android's descriptorless CelLiteRuntime. + CelLiteDescriptor fullDescriptor = TestAllTypesCelDescriptor.getDescriptor(); + MessageLiteDescriptor v1TestAllTypesDesc = + buildV1TestAllTypesDescriptor( + fullDescriptor + .getProtoTypeNamesToDescriptors() + .get(TestAllTypes.getDescriptor().getFullName())); + + ImmutableList.Builder allMsgDescs = ImmutableList.builder(); + for (MessageLiteDescriptor d : fullDescriptor.getProtoTypeNamesToDescriptors().values()) { + if (d.getProtoTypeName().equals(TestAllTypes.getDescriptor().getFullName())) { + allMsgDescs.add(v1TestAllTypesDesc); + } else { + allMsgDescs.add(d); + } + } + + CelLiteDescriptor v1Descriptor = new CelLiteDescriptor("v1", allMsgDescs.build()) {}; + ProtoMessageLiteValueProvider v1ValueProvider = + ProtoMessageLiteValueProvider.newInstance( + v1Descriptor, NestedTestAllTypesCelDescriptor.getDescriptor()); + + v1Runtime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .setStandardFunctions(CelStandardFunctions.ALL_STANDARD_FUNCTIONS) + .setValueProvider(v1ValueProvider) + .setContainer(CEL_CONTAINER) + .build(); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum UnknownScalarTestCase { + INT64("msg.single_int64", 0L, -42L), + UINT32("msg.single_uint32", UnsignedLong.ZERO, UnsignedLong.valueOf(123L)), + UINT64("msg.single_uint64", UnsignedLong.ZERO, UnsignedLong.valueOf(999L)), + SINT32("msg.single_sint32", 0L, -15L), + SINT64("msg.single_sint64", 0L, -250L), + FIXED32("msg.single_fixed32", UnsignedLong.ZERO, UnsignedLong.valueOf(320L)), + FIXED64("msg.single_fixed64", UnsignedLong.ZERO, UnsignedLong.valueOf(640L)), + SFIXED32("msg.single_sfixed32", 0L, -32L), + SFIXED64("msg.single_sfixed64", 0L, -64L), + FLOAT("msg.single_float", 0.0d, 1.5d), + DOUBLE("msg.single_double", 0.0d, 0.85d), + BOOL("msg.single_bool", false, true), + STRING("msg.single_string", "", "cel-skew-test"), + BYTES("msg.single_bytes", CelByteString.EMPTY, CelByteString.of("binary".getBytes(UTF_8))), + STANDALONE_ENUM("msg.standalone_enum", 0L, (long) NestedEnum.BAZ.getNumber()), + DURATION("msg.single_duration", Duration.ZERO, Duration.ofHours(1)), + TIMESTAMP("msg.single_timestamp", Instant.EPOCH, Instant.ofEpochSecond(1700000000L, 500L)), + ONEOF_BOOL("msg.oneof_bool", false, true), + SUBMESSAGE_SCALAR("msg.single_nested_message.bb", 0L, 123L); + + private final String expression; + private final Object expectedDefault; + private final Object expectedPopulated; + + UnknownScalarTestCase(String expression, Object expectedDefault, Object expectedPopulated) { + this.expression = expression; + this.expectedDefault = expectedDefault; + this.expectedPopulated = expectedPopulated; + } + } + + @Test + public void select_unsetUnknownScalar_returnsBakedDefault( + @TestParameter UnknownScalarTestCase testCase) throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval(testCase.expression, msg); + + assertThat(result).isEqualTo(testCase.expectedDefault); + } + + @Test + public void select_populatedUnknownScalar_decodesFromWireBytes( + @TestParameter UnknownScalarTestCase testCase) throws Exception { + Object result = eval(testCase.expression, POPULATED_V2_MESSAGE); + + assertThat(result).isEqualTo(testCase.expectedPopulated); + } + + @Test + public void select_populatedUnknownOneofEnum_decodesFromWireBytes() throws Exception { + TestAllTypes msg = TestAllTypes.newBuilder().setSingleNestedEnum(NestedEnum.BAR).build(); + + Object result = eval("msg.single_nested_enum == TestAllTypes.NestedEnum.BAR", msg); + + assertThat(result).isEqualTo(true); + } + + @Test + public void select_populatedUnknownOneofSubmessage_decodesFromWireBytes() throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder() + .setOneofType( + NestedTestAllTypes.newBuilder() + .setPayload(TestAllTypes.newBuilder().setSingleInt64(88L))) + .build(); + + Object result = eval("msg.oneof_type.payload.single_int64", msg); + + assertThat(result).isEqualTo(88L); + } + + @Test + public void select_unsetUnknownOneofSubmessage_returnsBakedDefault() throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval("msg.oneof_type.payload.single_int64", msg); + + assertThat(result).isEqualTo(0L); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum UnknownFieldPresenceTestCase { + INT64("has(msg.single_int64)"), + UINT64("has(msg.single_uint64)"), + DOUBLE("has(msg.single_double)"), + BOOL("has(msg.single_bool)"), + STRING("has(msg.single_string)"), + BYTES("has(msg.single_bytes)"), + STANDALONE_ENUM("has(msg.standalone_enum)"), + DURATION("has(msg.single_duration)"), + TIMESTAMP("has(msg.single_timestamp)"), + ONEOF_BOOL("has(msg.oneof_bool)"), + SUBMESSAGE("has(msg.single_nested_message)"), + SUBMESSAGE_SCALAR("has(msg.single_nested_message.bb)"), + REPEATED_INT64("has(msg.repeated_int64)"), + REPEATED_STRING("has(msg.repeated_string)"), + REPEATED_SUBMESSAGE("has(msg.repeated_nested_message)"), + MAP_INT32_INT32("has(msg.map_int32_int32)"); + + private final String expression; + + UnknownFieldPresenceTestCase(String expression) { + this.expression = expression; + } + } + + @Test + public void has_unknownField_whenPresentOnWire_returnsTrue( + @TestParameter UnknownFieldPresenceTestCase testCase) throws Exception { + Object result = eval(testCase.expression, POPULATED_V2_MESSAGE); + + assertThat(result).isEqualTo(true); + } + + @Test + public void has_unknownField_whenAbsentOnWire_returnsFalse( + @TestParameter UnknownFieldPresenceTestCase testCase) throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval(testCase.expression, msg); + + assertThat(result).isEqualTo(false); + } + + @Test + public void has_unknownSubmessageField_whenSubmessagePresentButFieldUnset_returnsFalse() + throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder() + .setSingleNestedMessage(NestedMessage.getDefaultInstance()) + .build(); + + Object result = eval("has(msg.single_nested_message.bb)", msg); + + assertThat(result).isEqualTo(false); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum UnknownRepeatedScalarTestCase { + REPEATED_INT64("msg.repeated_int64", ImmutableList.of(10L, 20L)), + REPEATED_DOUBLE("msg.repeated_double", ImmutableList.of(0.25d, 0.75d)), + REPEATED_STRING("msg.repeated_string", ImmutableList.of("foo", "bar")), + REPEATED_NESTED_ENUM( + "msg.repeated_nested_enum", + ImmutableList.of((long) NestedEnum.BAR.getNumber(), (long) NestedEnum.BAZ.getNumber())); + + private final String expression; + private final ImmutableList expectedElements; + + UnknownRepeatedScalarTestCase(String expression, ImmutableList expectedElements) { + this.expression = expression; + this.expectedElements = expectedElements; + } + } + + @Test + public void select_unsetUnknownRepeatedScalar_returnsBakedEmptyList( + @TestParameter UnknownRepeatedScalarTestCase testCase) throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval(testCase.expression, msg); + + assertThat((List) result).isEmpty(); + } + + @Test + public void select_populatedUnknownRepeatedScalar_decodesFromWireBytes( + @TestParameter UnknownRepeatedScalarTestCase testCase) throws Exception { + Object result = eval(testCase.expression, POPULATED_V2_MESSAGE); + + assertThat((List) result).containsExactlyElementsIn(testCase.expectedElements).inOrder(); + } + + @Test + public void select_populatedUnknownSubmessageLeaf_returnsRawMessage() throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder() + .setSingleNestedMessage(NestedMessage.newBuilder().setBb(123).build()) + .build(); + + Object result = eval("msg.single_nested_message", msg); + + assertThat(result).isInstanceOf(RawProtoMessageLiteValue.class); + } + + @Test + public void select_unsetUnknownMapField_returnsBakedDefault() throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval("msg.map_int32_int32", msg); + + assertThat(result).isEqualTo(ImmutableMap.of()); + } + + @Test + public void select_populatedUnknownMapField_throwsEvaluationException() { + TestAllTypes msg = TestAllTypes.newBuilder().putMapInt32Int32(1, 2).build(); + + CelEvaluationException thrown = + assertThrows(CelEvaluationException.class, () -> eval("msg.map_int32_int32", msg)); + + assertThat(thrown).hasCauseThat().isInstanceOf(UnsupportedOperationException.class); + assertThat(thrown).hasCauseThat().hasMessageThat().contains("map_int32_int32"); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum RenamedFieldTestCase { + SCALAR_INT32("msg.single_int32", "has(msg.single_int32)", 0L, 42L), + SUBMESSAGE_LEAF( + "msg.standalone_message", + "has(msg.standalone_message)", + NestedMessage.getDefaultInstance(), + NestedMessage.newBuilder().setBb(77).build()), + SUBMESSAGE_SCALAR("msg.standalone_message.bb", "has(msg.standalone_message.bb)", 0L, 77L), + REPEATED_INT32( + "msg.repeated_int32", + "has(msg.repeated_int32)", + ImmutableList.of(), + ImmutableList.of(10L, 20L)), + MAP_STRING_STRING( + "msg.map_string_string", + "has(msg.map_string_string)", + ImmutableMap.of(), + ImmutableMap.of("k", "v")), + MAP_INT64_MESSAGE( + "msg.map_int64_message", + "has(msg.map_int64_message)", + ImmutableMap.of(), + ImmutableMap.of(1L, NestedMessage.newBuilder().setBb(100).build())); + + private final String selectExpression; + private final String hasExpression; + private final Object expectedDefault; + private final Object expectedPopulated; + + RenamedFieldTestCase( + String selectExpression, + String hasExpression, + Object expectedDefault, + Object expectedPopulated) { + this.selectExpression = selectExpression; + this.hasExpression = hasExpression; + this.expectedDefault = expectedDefault; + this.expectedPopulated = expectedPopulated; + } + } + + @Test + public void select_renamedField_whenPopulated_resolvesByFieldNumber( + @TestParameter RenamedFieldTestCase testCase) throws Exception { + Object result = eval(testCase.selectExpression, POPULATED_RENAMED_MESSAGE); + + assertThat(result).isEqualTo(testCase.expectedPopulated); + } + + @Test + public void select_renamedField_whenUnset_returnsDefault( + @TestParameter RenamedFieldTestCase testCase) throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval(testCase.selectExpression, msg); + + assertThat(result).isEqualTo(testCase.expectedDefault); + } + + @Test + public void has_renamedField_whenPresent_returnsTrue(@TestParameter RenamedFieldTestCase testCase) + throws Exception { + Object result = eval(testCase.hasExpression, POPULATED_RENAMED_MESSAGE); + + assertThat(result).isEqualTo(true); + } + + @Test + public void has_renamedField_whenAbsent_returnsFalse(@TestParameter RenamedFieldTestCase testCase) + throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object result = eval(testCase.hasExpression, msg); + + assertThat(result).isEqualTo(false); + } + + @Test + public void has_renamedSubmessageField_whenSubmessagePresentButFieldUnset_returnsFalse() + throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder().setStandaloneMessage(NestedMessage.getDefaultInstance()).build(); + + Object result = eval("has(msg.standalone_message.bb)", msg); + + assertThat(result).isEqualTo(false); + } + + @Test + public void select_renamedMapField_indexing() throws Exception { + Object result = eval("msg.map_int64_message[1].bb", POPULATED_RENAMED_MESSAGE); + + assertThat(result).isEqualTo(100L); + } + + @Test + public void has_renamedRepeatedField_emptyPackedWireBytes_returnsFalse() throws Exception { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + CodedOutputStream cos = CodedOutputStream.newInstance(baos); + cos.writeBytes(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, ByteString.EMPTY); + cos.flush(); + TestAllTypes msg = + TestAllTypes.parseFrom(baos.toByteArray(), ExtensionRegistryLite.getEmptyRegistry()); + + Object result = eval("has(msg.repeated_int32)", msg); + + assertThat(result).isEqualTo(false); + } + + @Test + public void mixedExpression_versionSkewFieldsWithConditions() throws Exception { + Object result = + eval( + "msg.single_int64 < 0 && msg.single_bool && msg.single_string == 'cel-skew-test' &&" + + " msg.single_double > 0.7 && msg.standalone_enum == TestAllTypes.NestedEnum.BAZ" + + " && msg.single_duration == duration('1h') && msg.single_timestamp > timestamp(0)" + + " && has(msg.single_nested_message) && msg.single_nested_message.bb == 123", + POPULATED_V2_MESSAGE); + + assertThat(result).isEqualTo(true); + } + + @Test + public void mixedExpression_renamedMapFieldsWithCondition() throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder() + .putMapStringString("env", "prod") + .putMapInt64Message(42L, NestedMessage.newBuilder().setBb(99).build()) + .build(); + + Object result = + eval( + "has(msg.map_string_string) && msg.map_string_string['env'] == 'prod' &&" + + " msg.map_int64_message[42].bb == 99", + msg); + + assertThat(result).isEqualTo(true); + } + + @Test + public void fixedRootInput_populatedV1FieldsAndUnsetV2Fields_evaluatesV2RuleCleanly() + throws Exception { + // Models b/555872831#comment14: V1 APK constructs a fixed root payload with populated V1 + // fields while newer V2 fields (scalars, submessages, lists, maps, durations) are unset. + TestAllTypes v1Payload = + TestAllTypes.newBuilder() + .setOptionalString("us") + .setOptionalBool(true) + .addRepeatedUint64(10L) + .build(); + + Object result = + eval( + "has(msg.optional_string) && msg.optional_string == 'us' &&" + + " has(msg.optional_bool) && msg.optional_bool &&" + + " msg.repeated_uint64 == [10u] &&" + + " !has(msg.single_bool) && !msg.single_bool &&" + + " msg.single_int64 == 0 && msg.single_double < 0.7 &&" + + " msg.standalone_enum == TestAllTypes.NestedEnum.FOO &&" + + " msg.single_duration == duration('0s') &&" + + " msg.single_timestamp == timestamp(0) &&" + + " !has(msg.single_nested_message) && msg.single_nested_message.bb == 0 &&" + + " size(msg.repeated_nested_message) == 0 &&" + + " size(msg.map_int32_int32) == 0", + v1Payload); + + assertThat(result).isEqualTo(true); + } + + @Test + public void knownSubmessageChain_withUnsetV2SubfieldsInChildPayload_returnsBakedDefaults() + throws Exception { + // Models b/555872831#comment14: root NestedTestAllTypes and intermediate hops (child.payload) + // are known in V1, while schema skew occurs on V2 fields inside the child payload. + NestedTestAllTypes nestedMsg = + NestedTestAllTypes.newBuilder() + .setChild( + NestedTestAllTypes.newBuilder() + .setPayload( + TestAllTypes.newBuilder() + .setOptionalString("v1-child") + .setOptionalBool(true) + .build())) + .build(); + + Object result = + evalNested( + "nested_msg.child.payload.optional_string == 'v1-child' &&" + + " nested_msg.child.payload.optional_bool &&" + + " !has(nested_msg.child.payload.single_int64) &&" + + " nested_msg.child.payload.single_int64 == 0 &&" + + " !has(nested_msg.child.payload.single_nested_message.bb) &&" + + " nested_msg.child.payload.single_nested_message.bb == 0 &&" + + " size(nested_msg.child.payload.repeated_string) == 0", + nestedMsg); + + assertThat(result).isEqualTo(true); + } + + @Test + public void knownSubmessageChain_whenIntermediateChildIsUnset_returnsBakedDefaults() + throws Exception { + NestedTestAllTypes emptyNestedMsg = NestedTestAllTypes.getDefaultInstance(); + + Object result = + evalNested( + "!has(nested_msg.child.payload.single_int64) &&" + + " nested_msg.child.payload.single_int64 == 0 &&" + + " !has(nested_msg.child.payload.single_nested_message.bb) &&" + + " nested_msg.child.payload.single_nested_message.bb == 0 &&" + + " nested_msg.child.payload.single_duration == duration('0s')", + emptyNestedMsg); + + assertThat(result).isEqualTo(true); + } + + @Test + public void knownSubmessageChain_withPopulatedV2SubfieldsInChildPayload_decodesFromWireBytes() + throws Exception { + NestedTestAllTypes nestedMsg = + NestedTestAllTypes.newBuilder() + .setChild(NestedTestAllTypes.newBuilder().setPayload(POPULATED_V2_MESSAGE)) + .build(); + + Object result = + evalNested( + "has(nested_msg.child.payload.single_int64) &&" + + " nested_msg.child.payload.single_int64 == -42 &&" + + " has(nested_msg.child.payload.single_nested_message.bb) &&" + + " nested_msg.child.payload.single_nested_message.bb == 123 &&" + + " nested_msg.child.payload.single_duration == duration('1h') &&" + + " nested_msg.child.payload.repeated_string == ['foo', 'bar']", + nestedMsg); + + assertThat(result).isEqualTo(true); + } + + @Test + public void submessageDescriptorAbsentFromPool_whenUnsetAndPopulated_evaluatesViaRawBytes() + throws Exception { + // Simulates a V2 submessage type (NestedMessage) whose MessageLiteDescriptor is completely + // absent from V1's CelLiteDescriptorPool. + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + celOptimizer.optimize( + cel.compile( + "has(msg.single_nested_message.bb) ? msg.single_nested_message.bb :" + + " msg.single_nested_message.bb - 1") + .getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); + + Object unsetResult = program.eval(ImmutableMap.of("msg", TestAllTypes.getDefaultInstance())); + Object populatedResult = program.eval(ImmutableMap.of("msg", POPULATED_V2_MESSAGE)); + + assertThat(unsetResult).isEqualTo(-1L); + assertThat(populatedResult).isEqualTo(123L); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum ComprehensionTestCase { + EXISTS_MATCHING_SUBMESSAGE( + "msg.repeated_nested_message.exists(x, x.bb == 20)", POPULATED_V2_MESSAGE, true), + EXISTS_UNSET_SUBMESSAGE_LIST( + "msg.repeated_nested_message.exists(x, x.bb == 20)", + TestAllTypes.getDefaultInstance(), + false), + ALL_MATCHING_SUBMESSAGE( + "msg.repeated_nested_message.all(x, x.bb > 0)", POPULATED_V2_MESSAGE, true), + ALL_UNSET_SUBMESSAGE_LIST( + "msg.repeated_nested_message.all(x, x.bb > 0)", TestAllTypes.getDefaultInstance(), true), + EXISTS_ONE_SUBMESSAGE( + "msg.repeated_nested_message.exists_one(x, x.bb == 10)", POPULATED_V2_MESSAGE, true), + FILTER_AND_MAP_SUBMESSAGE( + "msg.repeated_nested_message.filter(x, x.bb > 10).map(x, x.bb * 2)", + POPULATED_V2_MESSAGE, + ImmutableList.of(40L)), + EXISTS_REPEATED_STRING("msg.repeated_string.exists(s, s == 'bar')", POPULATED_V2_MESSAGE, true); + + private final String expression; + private final TestAllTypes message; + private final Object expectedResult; + + ComprehensionTestCase(String expression, TestAllTypes message, Object expectedResult) { + this.expression = expression; + this.message = message; + this.expectedResult = expectedResult; + } + } + + @Test + public void comprehension_overUnknownRepeatedFields_evaluatesExpectedResult( + @TestParameter ComprehensionTestCase testCase) throws Exception { + Object result = eval(testCase.expression, testCase.message); + + assertThat(result).isEqualTo(testCase.expectedResult); + } + + @Test + public void nestedComprehension_acrossKnownAndUnknownRepeatedSubmessages_evaluatesCorrectly() + throws Exception { + NestedTestAllTypes nestedMsg = + NestedTestAllTypes.newBuilder() + .setPayload(TestAllTypes.newBuilder().addRepeatedInt32(5).addRepeatedInt32(20)) + .setChild(NestedTestAllTypes.newBuilder().setPayload(POPULATED_V2_MESSAGE)) + .build(); + + Object result = + evalNested( + "nested_msg.payload.repeated_int32.exists(x," + + " nested_msg.child.payload.repeated_nested_message.exists(y, x == y.bb))", + nestedMsg); + + assertThat(result).isEqualTo(true); + } + + @Test + public void celBind_unknownSubmessage_evaluatesBoundSubfields() throws Exception { + String expression = "cel.bind(sub, msg.single_nested_message, has(sub.bb) ? sub.bb : -1)"; + + Object populatedResult = eval(expression, POPULATED_V2_MESSAGE); + Object unsetResult = eval(expression, TestAllTypes.getDefaultInstance()); + + assertThat(populatedResult).isEqualTo(123L); + assertThat(unsetResult).isEqualTo(-1L); + } + + @Test + public void shortCircuiting_hasGuardPreventsEvaluatingPopulatedUnknownMap() throws Exception { + // msg.map_int32_int32 is populated on POPULATED_V2_MESSAGE and would throw + // UnsupportedOperationException if selected, but short-circuiting avoids evaluating it. + String expression = "!has(msg.map_int32_int32) ? size(msg.map_int32_int32) : 99"; + + Object populatedResult = eval(expression, POPULATED_V2_MESSAGE); + Object unsetResult = eval(expression, TestAllTypes.getDefaultInstance()); + + assertThat(populatedResult).isEqualTo(99L); + assertThat(unsetResult).isEqualTo(0L); + } + + @Test + public void has_multiHopUnknownSubmessage_whenUnset_returnsFalse() throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + + Object payloadResult = eval("has(msg.oneof_type.payload)", msg); + Object scalarResult = eval("has(msg.oneof_type.payload.single_int64)", msg); + + assertThat(payloadResult).isEqualTo(false); + assertThat(scalarResult).isEqualTo(false); + } + + @Test + public void has_multiHopUnknownSubmessage_whenIntermediateSetButLeafUnset_returnsFalse() + throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder().setOneofType(NestedTestAllTypes.getDefaultInstance()).build(); + + Object outerResult = eval("has(msg.oneof_type)", msg); + Object payloadResult = eval("has(msg.oneof_type.payload)", msg); + Object scalarResult = eval("has(msg.oneof_type.payload.single_int64)", msg); + + assertThat(outerResult).isEqualTo(true); + assertThat(payloadResult).isEqualTo(false); + assertThat(scalarResult).isEqualTo(false); + } + + @Test + public void + has_multiHopUnknownSubmessage_whenIntermediatePayloadSetButLeafScalarUnset_returnsFalse() + throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder() + .setOneofType( + NestedTestAllTypes.newBuilder().setPayload(TestAllTypes.getDefaultInstance())) + .build(); + + Object payloadResult = eval("has(msg.oneof_type.payload)", msg); + Object scalarResult = eval("has(msg.oneof_type.payload.single_int64)", msg); + + assertThat(payloadResult).isEqualTo(true); + assertThat(scalarResult).isEqualTo(false); + } + + @Test + public void has_multiHopUnknownSubmessage_whenAllHopsPopulated_returnsTrue() throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder() + .setOneofType( + NestedTestAllTypes.newBuilder() + .setPayload(TestAllTypes.newBuilder().setSingleInt64(42L))) + .build(); + + Object outerResult = eval("has(msg.oneof_type)", msg); + Object payloadResult = eval("has(msg.oneof_type.payload)", msg); + Object scalarResult = eval("has(msg.oneof_type.payload.single_int64)", msg); + Object valueResult = eval("msg.oneof_type.payload.single_int64", msg); + + assertThat(outerResult).isEqualTo(true); + assertThat(payloadResult).isEqualTo(true); + assertThat(scalarResult).isEqualTo(true); + assertThat(valueResult).isEqualTo(42L); + } + + @Test + public void oneof_activeDiscriminator_whenUnknownVariantPopulated_evaluatesPresenceAndFallback() + throws Exception { + // oneof_bool is unknown to V1 descriptor, oneof_msg is known to V1 descriptor. + TestAllTypes msg = TestAllTypes.newBuilder().setOneofBool(true).build(); + + Object unknownPresence = eval("has(msg.oneof_bool)", msg); + Object unknownValue = eval("msg.oneof_bool", msg); + Object knownPresence = eval("has(msg.oneof_msg)", msg); + Object knownScalarFallback = eval("msg.oneof_msg.bb", msg); + + assertThat(unknownPresence).isEqualTo(true); + assertThat(unknownValue).isEqualTo(true); + assertThat(knownPresence).isEqualTo(false); + assertThat(knownScalarFallback).isEqualTo(0L); + } + + @Test + public void oneof_activeDiscriminator_whenKnownVariantPopulated_evaluatesPresenceAndFallback() + throws Exception { + TestAllTypes msg = + TestAllTypes.newBuilder().setOneofMsg(NestedMessage.newBuilder().setBb(77).build()).build(); + + Object knownPresence = eval("has(msg.oneof_msg)", msg); + Object knownValue = eval("msg.oneof_msg.bb", msg); + Object unknownPresence = eval("has(msg.oneof_bool)", msg); + Object unknownValue = eval("msg.oneof_bool", msg); + + assertThat(knownPresence).isEqualTo(true); + assertThat(knownValue).isEqualTo(77L); + assertThat(unknownPresence).isEqualTo(false); + assertThat(unknownValue).isEqualTo(false); + } + + @Test + public void oneof_activeDiscriminator_whenExplicitFalseSetOnUnknownVariant_evaluatesPresenceTrue() + throws Exception { + TestAllTypes msg = TestAllTypes.newBuilder().setOneofBool(false).build(); + + Object unknownPresence = eval("has(msg.oneof_bool)", msg); + Object unknownValue = eval("msg.oneof_bool", msg); + + assertThat(unknownPresence).isEqualTo(true); + assertThat(unknownValue).isEqualTo(false); + } + + @Test + public void select_unknownEnumValueBeyondEnumRange_decodesRawNumericValue() throws Exception { + TestAllTypes msg = TestAllTypes.newBuilder().setStandaloneEnumValue(99).build(); + + Object presenceResult = eval("has(msg.standalone_enum)", msg); + Object valueResult = eval("msg.standalone_enum", msg); + Object comparisonResult = eval("msg.standalone_enum == 99", msg); + + assertThat(presenceResult).isEqualTo(true); + assertThat(valueResult).isEqualTo(99L); + assertThat(comparisonResult).isEqualTo(true); + } + + @Test + public void select_negativeInt32InPopulatedMessage_evaluatesCorrectly() throws Exception { + TestAllTypes msg = TestAllTypes.newBuilder().setSingleInt32(-42).build(); + + Object valueResult = eval("msg.single_int32", msg); + Object comparisonResult = eval("msg.single_int32 == -42 && msg.single_int32 < 0", msg); + + assertThat(valueResult).isEqualTo(-42L); + assertThat(comparisonResult).isEqualTo(true); + } + + @Test + public void submessageEquality_reflexiveAndIdenticalMessages_evaluatesTrue() throws Exception { + Object selfEqualResult = + eval("msg.single_nested_message == msg.single_nested_message", POPULATED_V2_MESSAGE); + Object emptyEqualResult = + eval( + "msg.single_nested_message == msg.single_nested_message", + TestAllTypes.getDefaultInstance()); + + assertThat(selfEqualResult).isEqualTo(true); + assertThat(emptyEqualResult).isEqualTo(true); + } + + @Test + public void submessageEquality_differentSubmessages_evaluatesFalse() throws Exception { + Object differentResult = + eval("msg.single_nested_message == msg.repeated_nested_message[0]", POPULATED_V2_MESSAGE); + + assertThat(differentResult).isEqualTo(false); + } + + @Test + public void select_unknownNegativeInt32_decodesFromRawWireBytes() throws Exception { + CelLiteRuntime runtimeWithoutInt32 = newRuntimeWithoutSingleInt32Descriptor(); + CelAbstractSyntaxTree ast = cel.compile("msg.single_int32").getAst(); + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + Program program = runtimeWithoutInt32.createProgram(optimizedAst); + TestAllTypes msg = TestAllTypes.newBuilder().setSingleInt32(-42).build(); + + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(result).isEqualTo(-42L); + } + + @Test + public void select_populatedUnknownSignedNumbers_decodesNegativeValues() throws Exception { + Object sint32Result = eval("msg.single_sint32", POPULATED_V2_MESSAGE); + Object sint64Result = eval("msg.single_sint64", POPULATED_V2_MESSAGE); + Object sfixed32Result = eval("msg.single_sfixed32", POPULATED_V2_MESSAGE); + Object sfixed64Result = eval("msg.single_sfixed64", POPULATED_V2_MESSAGE); + + assertThat(sint32Result).isEqualTo(-15L); + assertThat(sint64Result).isEqualTo(-250L); + assertThat(sfixed32Result).isEqualTo(-32L); + assertThat(sfixed64Result).isEqualTo(-64L); + } + + @Test + public void checkedExprWireRoundTrip_preservesOptimizedSelectsUnderVersionSkew() + throws Exception { + CelAbstractSyntaxTree ast = + cel.compile( + "has(msg.single_nested_message.bb) && msg.single_nested_message.bb == 123 &&" + + " msg.single_duration == duration('1h') &&" + + " msg.repeated_string.exists(s, s == 'foo')") + .getAst(); + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + byte[] serializedCheckedExpr = + CelProtoV1Alpha1AbstractSyntaxTree.fromCelAst(optimizedAst).toCheckedExpr().toByteArray(); + CheckedExpr deserializedCheckedExpr = + CheckedExpr.parseFrom(serializedCheckedExpr, ExtensionRegistryLite.getEmptyRegistry()); + CelAbstractSyntaxTree deserializedAst = + CelProtoV1Alpha1AbstractSyntaxTree.fromCheckedExpr(deserializedCheckedExpr).getAst(); + + Object result = + v1Runtime.createProgram(deserializedAst).eval(ImmutableMap.of("msg", POPULATED_V2_MESSAGE)); + + assertThat(result).isEqualTo(true); + } + + @Test + public void select_parsedOnlyAst_dispatchesDynamicallyByFunctionName() throws Exception { + TestAllTypes msg = TestAllTypes.newBuilder().setSingleInt64(99L).build(); + CelAbstractSyntaxTree ast = cel.compile("msg.single_int64").getAst(); + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + CelAbstractSyntaxTree parsedOptimizedAst = + CelAbstractSyntaxTree.newParsedAst(optimizedAst.getExpr(), optimizedAst.getSource()); + + Program program = v1Runtime.createProgram(parsedOptimizedAst); + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(parsedOptimizedAst.isChecked()).isFalse(); + assertThat(result).isEqualTo(99L); + } + + @Test + public void has_parsedOnlyAst_dispatchesDynamicallyByFunctionName() throws Exception { + TestAllTypes msg = TestAllTypes.newBuilder().setSingleInt64(99L).build(); + CelAbstractSyntaxTree ast = cel.compile("has(msg.single_int64)").getAst(); + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + CelAbstractSyntaxTree parsedOptimizedAst = + CelAbstractSyntaxTree.newParsedAst(optimizedAst.getExpr(), optimizedAst.getSource()); + + Program program = v1Runtime.createProgram(parsedOptimizedAst); + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(parsedOptimizedAst.isChecked()).isFalse(); + assertThat(result).isEqualTo(true); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum CelBlockOptimizationTestCase { + REPEATED_UNKNOWN_SCALAR_MATCHING( + "msg.single_int64 < -10 && msg.single_int64 > -100", POPULATED_V2_MESSAGE, true), + REPEATED_UNKNOWN_SCALAR_NON_MATCHING( + "msg.single_int64 < -10 && msg.single_int64 > -100", + TestAllTypes.getDefaultInstance(), + false), + REPEATED_UNKNOWN_SUBMESSAGE_MATCHING( + "msg.single_nested_message.bb > 10 && msg.single_nested_message.bb < 200", + POPULATED_V2_MESSAGE, + true), + REPEATED_UNKNOWN_SUBMESSAGE_NON_MATCHING( + "msg.single_nested_message.bb > 10 && msg.single_nested_message.bb < 200", + TestAllTypes.getDefaultInstance(), + false), + REPEATED_UNKNOWN_HAS_FIELD_MATCHING( + "has(msg.single_nested_message.bb) && (has(msg.single_nested_message.bb) ||" + + " msg.single_int64 > 0)", + POPULATED_V2_MESSAGE, + true), + REPEATED_UNKNOWN_HAS_FIELD_NON_MATCHING( + "has(msg.single_nested_message.bb) && (has(msg.single_nested_message.bb) ||" + + " msg.single_int64 > 0)", + TestAllTypes.getDefaultInstance(), + false), + REPEATED_UNKNOWN_DURATION_MATCHING( + "msg.single_duration > duration('1s') && msg.single_duration < duration('2h')", + POPULATED_V2_MESSAGE, + true), + REPEATED_UNKNOWN_DURATION_NON_MATCHING( + "msg.single_duration > duration('1s') && msg.single_duration < duration('2h')", + TestAllTypes.getDefaultInstance(), + false); + + private final String expression; + private final TestAllTypes message; + private final boolean expectedResult; + + CelBlockOptimizationTestCase(String expression, TestAllTypes message, boolean expectedResult) { + this.expression = expression; + this.message = message; + this.expectedResult = expectedResult; + } + } + + @Test + public void celBlock_subexpressionThenSelectOptimizer_evaluatesExpectedResult( + @TestParameter CelBlockOptimizationTestCase testCase) throws Exception { + CelOptimizer blockOptimizer = newSubexpressionThenSelectOptimizer(); + CelAbstractSyntaxTree ast = cel.compile(testCase.expression).getAst(); + + CelAbstractSyntaxTree optimizedAst = blockOptimizer.optimize(ast); + Object result = + v1Runtime.createProgram(optimizedAst).eval(ImmutableMap.of("msg", testCase.message)); + + assertThat(CelBlock.extract(optimizedAst)).isPresent(); + assertThat(result).isEqualTo(testCase.expectedResult); + } + + @Test + public void celBlock_selectThenSubexpressionOptimizer_evaluatesExpectedResult( + @TestParameter CelBlockOptimizationTestCase testCase) throws Exception { + CelOptimizer blockOptimizer = newSelectThenSubexpressionOptimizer(); + CelAbstractSyntaxTree ast = cel.compile(testCase.expression).getAst(); + + CelAbstractSyntaxTree optimizedAst = blockOptimizer.optimize(ast); + Object result = + v1Runtime.createProgram(optimizedAst).eval(ImmutableMap.of("msg", testCase.message)); + + assertThat(CelBlock.extract(optimizedAst)).isPresent(); + assertThat(result).isEqualTo(testCase.expectedResult); + } + + @Test + public void celBlock_subexpressionThenSelectOptimizer_sharedUnknownSubmessageAcrossHasAndSelect() + throws Exception { + CelOptimizer blockOptimizer = newSubexpressionThenSelectOptimizer(); + CelAbstractSyntaxTree ast = + cel.compile("has(msg.single_nested_message.bb) && msg.single_nested_message.bb == 123") + .getAst(); + CelAbstractSyntaxTree optimizedAst = blockOptimizer.optimize(ast); + Program program = v1Runtime.createProgram(optimizedAst); + + Object populatedResult = program.eval(ImmutableMap.of("msg", POPULATED_V2_MESSAGE)); + Object unsetResult = program.eval(ImmutableMap.of("msg", TestAllTypes.getDefaultInstance())); + + assertThat(CelBlock.extract(optimizedAst)).isPresent(); + assertThat(populatedResult).isEqualTo(true); + assertThat(unsetResult).isEqualTo(false); + } + + private Object eval(String expression, TestAllTypes message) throws Exception { + CelAbstractSyntaxTree ast = cel.compile(expression).getAst(); + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + Program program = v1Runtime.createProgram(optimizedAst); + return program.eval(ImmutableMap.of("msg", message)); + } + + private Object evalNested(String expression, NestedTestAllTypes nestedMessage) throws Exception { + CelAbstractSyntaxTree ast = cel.compile(expression).getAst(); + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + Program program = v1Runtime.createProgram(optimizedAst); + return program.eval(ImmutableMap.of("nested_msg", nestedMessage)); + } + + private static MessageLiteDescriptor buildV1TestAllTypesDescriptor( + MessageLiteDescriptor fullMsgDesc) { + return buildV1TestAllTypesDescriptor(fullMsgDesc, /* additionalExcludedField= */ null); + } + + private static MessageLiteDescriptor buildV1TestAllTypesDescriptor( + MessageLiteDescriptor fullMsgDesc, @Nullable String additionalExcludedField) { + ImmutableList.Builder v1Fields = ImmutableList.builder(); + for (FieldLiteDescriptor f : fullMsgDesc.getFieldDescriptors()) { + if (SKEW_EXCLUDED_FIELD_NAMES.contains(f.getFieldName()) + || f.getFieldName().equals(additionalExcludedField)) { + continue; + } + String v1FieldName = + SKEW_RENAMED_FIELD_NAMES.getOrDefault(f.getFieldName(), f.getFieldName()); + if (!v1FieldName.equals(f.getFieldName())) { + v1Fields.add( + new FieldLiteDescriptor( + f.getFieldNumber(), + v1FieldName, + f.getJavaType(), + f.getEncodingType(), + f.getProtoFieldType(), + f.getIsPacked(), + f.getFieldProtoTypeName())); + } else { + v1Fields.add(f); + } + } + return new MessageLiteDescriptor( + fullMsgDesc.getProtoTypeName(), v1Fields.build(), fullMsgDesc::newMessageBuilder); + } + + private static CelLiteRuntime newRuntimeWithoutNestedMessageDescriptor() { + CelLiteDescriptor fullDescriptor = TestAllTypesCelDescriptor.getDescriptor(); + MessageLiteDescriptor v1TestAllTypesDesc = + buildV1TestAllTypesDescriptor( + fullDescriptor + .getProtoTypeNamesToDescriptors() + .get(TestAllTypes.getDescriptor().getFullName())); + CelLiteDescriptor partialDescriptor = + new CelLiteDescriptor("v1_no_nested", ImmutableList.of(v1TestAllTypesDesc)) {}; + return CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .setStandardFunctions(CelStandardFunctions.ALL_STANDARD_FUNCTIONS) + .setValueProvider(ProtoMessageLiteValueProvider.newInstance(partialDescriptor)) + .setContainer(CEL_CONTAINER) + .build(); + } + + private static CelLiteRuntime newRuntimeWithoutSingleInt32Descriptor() { + CelLiteDescriptor fullDescriptor = TestAllTypesCelDescriptor.getDescriptor(); + MessageLiteDescriptor v1TestAllTypesDesc = + buildV1TestAllTypesDescriptor( + fullDescriptor + .getProtoTypeNamesToDescriptors() + .get(TestAllTypes.getDescriptor().getFullName()), + "single_int32"); + CelLiteDescriptor partialDescriptor = + new CelLiteDescriptor("v1_no_int32", ImmutableList.of(v1TestAllTypesDesc)) {}; + return CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .setStandardFunctions(CelStandardFunctions.ALL_STANDARD_FUNCTIONS) + .setValueProvider(ProtoMessageLiteValueProvider.newInstance(partialDescriptor)) + .setContainer(CEL_CONTAINER) + .build(); + } + + private CelOptimizer newSubexpressionThenSelectOptimizer() { + return CelOptimizerFactory.standardCelOptimizerBuilder(cel) + .addAstOptimizers( + SubexpressionOptimizer.getInstance(), + SelectOptimizer.newInstance( + SelectOptimizerOptions.newBuilder().build(), + TestAllTypes.getDescriptor().getFile())) + .build(); + } + + private CelOptimizer newSelectThenSubexpressionOptimizer() { + return CelOptimizerFactory.standardCelOptimizerBuilder(cel) + .addAstOptimizers( + SelectOptimizer.newInstance( + SelectOptimizerOptions.newBuilder().build(), + TestAllTypes.getDescriptor().getFile()), + SubexpressionOptimizer.newInstance( + SubexpressionOptimizerOptions.newBuilder() + .addEliminableFunctions("cel.@attribute", "cel.@hasField") + .build())) + .build(); + } +}