Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 20 additions & 2 deletions checker/src/main/java/dev/cel/checker/ExprChecker.java
Original file line number Diff line number Diff line change
Expand Up @@ -266,7 +266,7 @@ private void visit(CelMutableExpr expr, CelMutableIdent ident) {
// Overwrite the identifier with its fully qualified name.
expr.setIdent(CelMutableIdent.create(refName));
}
env.setType(expr, decl.type());
env.setType(expr, normalizeIdentType(decl.type()));
env.setRef(expr, makeReference(refName, decl));
}

Expand All @@ -287,7 +287,7 @@ private void visit(CelMutableExpr expr, CelMutableSelect select) {
// variable name.
expr.setIdent(CelMutableIdent.create(refName));
}
env.setType(expr, decl.type());
env.setType(expr, normalizeIdentType(decl.type()));

env.setRef(expr, makeReference(refName, decl));
}
Expand Down Expand Up @@ -855,6 +855,24 @@ private static CelType normalizeFieldType(CelType celType) {
return celType;
}

private static CelType normalizeIdentType(CelType type) {
if (type instanceof TypeType) {
TypeType typeType = (TypeType) type;
CelType typeOfType = typeType.type();
if (typeOfType.kind() == CelKind.STRUCT) {
switch (typeOfType.name()) {
case CelTypes.DURATION_MESSAGE:
return TypeType.create(SimpleType.DURATION);
case CelTypes.TIMESTAMP_MESSAGE:
return TypeType.create(SimpleType.TIMESTAMP);
default:
break;
}
}
}
return type;
}

/** TODO: Remove after cl/984117942 is submitted. */
private static Optional<CelType> lookupLegacyFieldType(
TypeProvider legacyTypeProvider, CelType type, String fieldName) {
Expand Down
4 changes: 4 additions & 0 deletions checker/src/test/java/dev/cel/checker/ExprCheckerTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
78 changes: 76 additions & 2 deletions checker/src/test/java/dev/cel/checker/TypesTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,12 @@

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.types.CelKind;
Expand All @@ -37,9 +42,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
Expand Down Expand Up @@ -350,6 +354,76 @@ 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));

final String expression;
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);

final String expression;
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_typeParamInCompositeTypeType_resolvesReturnType() throws Exception {
TypeParamType typeParamT = TypeParamType.create("T");
Expand Down
23 changes: 22 additions & 1 deletion checker/src/test/resources/types.baseline
Original file line number Diff line number Diff line change
Expand Up @@ -45,4 +45,25 @@ __comprehension__(
]~list(list(dyn))
)~list(list(dyn))^add_list,
// Result
@result~list(list(dyn))^@result)~list(list(dyn))
@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
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -1260,4 +1261,88 @@ public void selectByFieldNumber_absentMessageFieldWithoutDescriptor_returnsUnkno
.hasMessageThat()
.contains("Decoding unknown map field from wire bytes is unsupported");
}

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);

final TestAllTypes populatedProto;
final SelectField selectField;
final Object expectedPopulatedValue;
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<Object> populatedFound = populatedRaw.findByFieldNumber(testCase.selectField);
Optional<Object> emptyFound = emptyRaw.findByFieldNumber(testCase.selectField);

assertThat(populatedFound).hasValue(testCase.expectedPopulatedValue);
assertThat(emptyFound).isEmpty();
}
}
Loading
Loading