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/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(