diff --git a/checker/BUILD.bazel b/checker/BUILD.bazel index ac00ddff1..9beedf8e9 100644 --- a/checker/BUILD.bazel +++ b/checker/BUILD.bazel @@ -1,4 +1,5 @@ load("@rules_java//java:defs.bzl", "java_library") +load("//:cel_android_rules.bzl", "cel_android_library") package( default_applicable_licenses = ["//:license"], @@ -53,3 +54,8 @@ java_library( name = "standard_decl", exports = ["//checker/src/main/java/dev/cel/checker:standard_decl"], ) + +cel_android_library( + name = "standard_decl_android", + exports = ["//checker/src/main/java/dev/cel/checker:standard_decl_android"], +) diff --git a/checker/src/main/java/dev/cel/checker/BUILD.bazel b/checker/src/main/java/dev/cel/checker/BUILD.bazel index 304ce0ec4..c42083220 100644 --- a/checker/src/main/java/dev/cel/checker/BUILD.bazel +++ b/checker/src/main/java/dev/cel/checker/BUILD.bazel @@ -1,4 +1,5 @@ load("@rules_java//java:defs.bzl", "java_library") +load("//:cel_android_rules.bzl", "cel_android_library") package( default_applicable_licenses = [ @@ -69,7 +70,6 @@ java_library( ":checker_legacy_environment", ":proto_type_mask", ":standard_decl", - ":type_provider_legacy_impl", "//common:cel_ast", "//common:cel_descriptor_util", "//common:cel_function_decl", @@ -235,3 +235,23 @@ java_library( "@maven//:com_google_guava_guava", ], ) + +cel_android_library( + name = "standard_decl_android", + srcs = [ + "CelStandardDeclarations.java", + ], + tags = [ + ], + deps = [ + "//common:cel_function_decl_android", + "//common:cel_overload_decl_android", + "//common:cel_var_decl_android", + "//common:operator_android", + "//common/types:cel_types_android", + "//common/types:type_providers_android", + "//common/types:types_android", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven_android//:com_google_guava_guava", + ], +) diff --git a/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java b/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java index 3384ccd63..929050fef 100644 --- a/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java +++ b/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java @@ -111,6 +111,8 @@ public CelTypeProvider getTypeProvider() { } @Override + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") public CelCheckerBuilder toCheckerBuilder() { CelCheckerBuilder builder = new Builder() @@ -124,6 +126,10 @@ public CelCheckerBuilder toCheckerBuilder() { .addFileTypes(fileDescriptors) .addProtoTypeMasks(protoTypeMasks); + if (typeProvider != null) { + builder.setTypeProvider(typeProvider); + } + if (expectedResultType.isPresent()) { builder.setResultType(expectedResultType.get()); } @@ -162,11 +168,13 @@ public void accept(EnvVisitor envVisitor) { private Env getEnv(Errors errors) { Env env; if (overriddenStandardDeclarations != null) { - env = Env.standard(overriddenStandardDeclarations, errors, typeProvider, celOptions); + env = + Env.standard( + overriddenStandardDeclarations, errors, celTypeProvider, typeProvider, celOptions); } else if (standardEnvironmentEnabled) { - env = Env.standard(errors, typeProvider, celOptions); + env = Env.standard(errors, celTypeProvider, typeProvider, celOptions); } else { - env = Env.unconfigured(errors, typeProvider, celOptions); + env = Env.unconfigured(errors, celTypeProvider, typeProvider, celOptions); } identDeclarations.forEach(env::add); functionDeclarations.forEach(env::add); @@ -459,20 +467,13 @@ public CelCheckerLegacyImpl build() { messageTypeProvider = protoTypeMaskTypeProvider; } - TypeProvider legacyProvider = new TypeProviderLegacyImpl(messageTypeProvider); - if (customTypeProvider != null) { - legacyProvider = - new TypeProvider.CombinedTypeProvider( - ImmutableList.of(customTypeProvider, legacyProvider)); - } - return new CelCheckerLegacyImpl( celOptions, container, identDeclarationSet, functionDeclarations.build(), Optional.fromNullable(expectedResultType), - legacyProvider, + customTypeProvider, messageTypeProvider, standardEnvironmentEnabled, standardDeclarations, diff --git a/checker/src/main/java/dev/cel/checker/Env.java b/checker/src/main/java/dev/cel/checker/Env.java index 81fec362c..e2d18ae4e 100644 --- a/checker/src/main/java/dev/cel/checker/Env.java +++ b/checker/src/main/java/dev/cel/checker/Env.java @@ -46,8 +46,14 @@ import dev.cel.common.types.CelKind; import dev.cel.common.types.CelProtoTypes; import dev.cel.common.types.CelType; +import dev.cel.common.types.CelTypeProvider; import dev.cel.common.types.CelTypes; +import dev.cel.common.types.EnumType; +import dev.cel.common.types.OpaqueType; import dev.cel.common.types.SimpleType; +import dev.cel.common.types.StructType; +import dev.cel.common.types.StructTypeReference; +import dev.cel.common.types.TypeType; import dev.cel.parser.CelStandardMacro; import java.util.ArrayList; import java.util.HashMap; @@ -80,8 +86,15 @@ public class Env { public static final CelFunctionDecl ERROR_FUNCTION_DECL = CelFunctionDecl.newBuilder().setName("*error*").build(); - /** Type provider responsible for resolving CEL message references to strong types. */ - private final TypeProvider typeProvider; + private static final CelTypeProvider EMPTY_TYPE_PROVIDER = + new CelTypeProvider.CombinedCelTypeProvider(ImmutableList.of()); + + /** Type provider responsible for resolving CEL types. */ + private final CelTypeProvider celTypeProvider; + + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + private final @Nullable TypeProvider legacyTypeProvider; /** * Stack of declaration groups where each entry in stack represents a scope capable of hinding @@ -107,14 +120,6 @@ public class Env { .enableNamespacedDeclarations(false) .build(); - private Env( - Errors errors, TypeProvider typeProvider, DeclGroup declGroup, CelOptions celOptions) { - this.celOptions = celOptions; - this.errors = Preconditions.checkNotNull(errors); - this.typeProvider = Preconditions.checkNotNull(typeProvider); - this.decls.add(Preconditions.checkNotNull(declGroup)); - } - /** * @deprecated Do not use. This exists for compatibility reasons. Migrate to CEL-Java fluent APIs. * See {@code CelCompilerFactory}. @@ -133,6 +138,24 @@ static Env unconfigured(Errors errors, CelOptions celOptions) { return unconfigured(errors, new DescriptorTypeProvider(), celOptions); } + /** + * Creates an unconfigured {@code Env} value without the standard CEL types, functions, and + * operators using a custom {@code celTypeProvider}. + */ + static Env unconfigured(Errors errors, CelTypeProvider celTypeProvider, CelOptions celOptions) { + return unconfigured(errors, celTypeProvider, null, celOptions); + } + + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + static Env unconfigured( + Errors errors, + CelTypeProvider celTypeProvider, + @Nullable TypeProvider legacyTypeProvider, + CelOptions celOptions) { + return new Env(errors, celTypeProvider, legacyTypeProvider, new DeclGroup(), celOptions); + } + /** * Creates an unconfigured {@code Env} value without the standard CEL types, functions, and * operators using a custom {@code typeProvider}. @@ -142,7 +165,7 @@ static Env unconfigured(Errors errors, CelOptions celOptions) { */ @Deprecated public static Env unconfigured(Errors errors, TypeProvider typeProvider, CelOptions celOptions) { - return new Env(errors, typeProvider, new DeclGroup(), celOptions); + return unconfigured(errors, EMPTY_TYPE_PROVIDER, typeProvider, celOptions); } /** @@ -163,6 +186,22 @@ public static Env standard(Errors errors, TypeProvider typeProvider) { return standard(errors, typeProvider, LEGACY_TYPE_CHECKER_OPTIONS); } + static Env standard(Errors errors, CelTypeProvider celTypeProvider, CelOptions celOptions) { + return standard(errors, celTypeProvider, null, celOptions); + } + + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + static Env standard( + Errors errors, + CelTypeProvider celTypeProvider, + @Nullable TypeProvider legacyTypeProvider, + CelOptions celOptions) { + CelStandardDeclarations celStandardDeclaration = newStandardDeclarations(celOptions); + return standard( + celStandardDeclaration, errors, celTypeProvider, legacyTypeProvider, celOptions); + } + /** * Creates an {@code Env} value configured with the standard types, functions, and operators, * configured with a custom {@code typeProvider} and a reference to the {@code celOptions} to use @@ -176,48 +215,27 @@ public static Env standard(Errors errors, TypeProvider typeProvider) { */ @Deprecated public static Env standard(Errors errors, TypeProvider typeProvider, CelOptions celOptions) { - CelStandardDeclarations celStandardDeclaration = - CelStandardDeclarations.newBuilder() - .filterFunctions( - (function, overload) -> { - switch (function) { - case INT: - if (!celOptions.enableUnsignedLongs() - && overload.equals(Conversions.INT64_TO_INT64)) { - return false; - } - break; - case TIMESTAMP: - // TODO: Remove this flag guard once the feature has been - // auto-enabled. - if (!celOptions.enableTimestampEpoch() - && overload.equals(Conversions.INT64_TO_TIMESTAMP)) { - return false; - } - break; - default: - if (!celOptions.enableHeterogeneousNumericComparisons() - && overload instanceof Comparison) { - Comparison comparison = (Comparison) overload; - if (comparison.isHeterogeneousComparison()) { - return false; - } - } - break; - } - return true; - }) - .build(); + return standard( + newStandardDeclarations(celOptions), errors, EMPTY_TYPE_PROVIDER, typeProvider, celOptions); + } - return standard(celStandardDeclaration, errors, typeProvider, celOptions); + static Env standard( + CelStandardDeclarations celStandardDeclaration, + Errors errors, + CelTypeProvider celTypeProvider, + CelOptions celOptions) { + return standard(celStandardDeclaration, errors, celTypeProvider, null, celOptions); } - public static Env standard( + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + static Env standard( CelStandardDeclarations celStandardDeclaration, Errors errors, - TypeProvider typeProvider, + CelTypeProvider celTypeProvider, + @Nullable TypeProvider legacyTypeProvider, CelOptions celOptions) { - Env env = Env.unconfigured(errors, typeProvider, celOptions); + Env env = Env.unconfigured(errors, celTypeProvider, legacyTypeProvider, celOptions); // Isolate the standard declarations into their own scope for forward compatibility. celStandardDeclaration.functionDecls().forEach(env::add); celStandardDeclaration.identifierDecls().forEach(env::add); @@ -226,14 +244,70 @@ public static Env standard( return env; } + @Deprecated + public static Env standard( + CelStandardDeclarations celStandardDeclaration, + Errors errors, + TypeProvider typeProvider, + CelOptions celOptions) { + return standard(celStandardDeclaration, errors, EMPTY_TYPE_PROVIDER, typeProvider, celOptions); + } + + private static CelStandardDeclarations newStandardDeclarations(CelOptions celOptions) { + return CelStandardDeclarations.newBuilder() + .filterFunctions( + (function, overload) -> { + switch (function) { + case INT: + if (!celOptions.enableUnsignedLongs() + && overload.equals(Conversions.INT64_TO_INT64)) { + return false; + } + break; + case TIMESTAMP: + // TODO: Remove this flag guard once the feature has been + // auto-enabled. + if (!celOptions.enableTimestampEpoch() + && overload.equals(Conversions.INT64_TO_TIMESTAMP)) { + return false; + } + break; + default: + if (!celOptions.enableHeterogeneousNumericComparisons() + && overload instanceof Comparison) { + Comparison comparison = (Comparison) overload; + if (comparison.isHeterogeneousComparison()) { + return false; + } + } + break; + } + return true; + }) + .build(); + } + /** Returns the current Errors object. */ public Errors getErrorContext() { return errors; } - /** Returns the {@code TypeProvider}. */ - public TypeProvider getTypeProvider() { - return typeProvider; + /** Returns the {@code CelTypeProvider}. */ + public CelTypeProvider getCelTypeProvider() { + return celTypeProvider; + } + + /** + * Returns the {@code TypeProvider}, or {@code null} if only modern {@link CelTypeProvider} was + * configured. + * + * @deprecated Use {@link #getCelTypeProvider()} instead. + */ + @Deprecated + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + @Nullable TypeProvider getTypeProvider() { + return legacyTypeProvider; } /** @@ -478,7 +552,14 @@ public Env add(String name, Type type) { // Next try to import the name as a reference to a message type. // This is done via the type provider. - Optional type = typeProvider.lookupCelType(cand); + Optional type = + celTypeProvider + .findType(cand) + .filter(t -> !(t instanceof EnumType) || t.name().equals(cand)) + .map(Env::wrapAsTypeIfNeeded); + if (!type.isPresent() && legacyTypeProvider != null) { + type = legacyTypeProvider.lookupCelType(cand); + } if (type.isPresent()) { decl = CelVarDecl.newVarDeclaration(cand, type.get()); decls.get(0).putIdent(decl); @@ -487,13 +568,13 @@ public Env add(String name, Type type) { // Next try to import this as an enum value by splitting the name in a type prefix and // the enum inside. - Integer enumValue = typeProvider.lookupEnumValue(cand); - if (enumValue != null) { + Optional enumValue = lookupEnumValue(cand); + if (enumValue.isPresent()) { decl = CelVarDecl.newBuilder() .setName(cand) .setType(SimpleType.INT) - .setConstant(CelConstant.ofValue(enumValue)) + .setConstant(CelConstant.ofValue(enumValue.get())) .build(); decls.get(0).putIdent(decl); @@ -502,6 +583,44 @@ public Env add(String name, Type type) { return null; } + private Optional lookupEnumValue(String enumName) { + int dotIndex = enumName.lastIndexOf("."); + if (dotIndex > 0 && dotIndex < enumName.length() - 1) { + String enumTypeName = enumName.substring(0, dotIndex); + String localEnumName = enumName.substring(dotIndex + 1); + Optional enumValue = + celTypeProvider + .findType(enumTypeName) + .filter(t -> t instanceof EnumType) + .flatMap(t -> ((EnumType) t).findNumberByName(localEnumName)); + if (enumValue.isPresent()) { + return enumValue; + } + enumValue = + celTypeProvider + .findType(enumName) + .filter(t -> t instanceof EnumType) + .flatMap(t -> ((EnumType) t).findNumberByName(localEnumName)); + if (enumValue.isPresent()) { + return enumValue; + } + } + if (legacyTypeProvider != null) { + return Optional.ofNullable(legacyTypeProvider.lookupEnumValue(enumName)); + } + return Optional.empty(); + } + + private static CelType wrapAsTypeIfNeeded(CelType type) { + if (type instanceof StructType + || type instanceof StructTypeReference + || type instanceof EnumType + || type instanceof OpaqueType) { + return TypeType.create(type); + } + return type; + } + /** * Lookup a local identifier by name. This searches only comprehension scopes, bypassing standard * environment or user-defined environment. @@ -1015,4 +1134,19 @@ static CelType getWellKnownType(CelType type) { Preconditions.checkArgument(type.kind() == CelKind.STRUCT); return CelTypes.getWellKnownCelType(type.name()).get(); } + + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + private Env( + Errors errors, + CelTypeProvider celTypeProvider, + @Nullable TypeProvider legacyTypeProvider, + DeclGroup declGroup, + CelOptions celOptions) { + this.celOptions = Preconditions.checkNotNull(celOptions); + this.errors = Preconditions.checkNotNull(errors); + this.celTypeProvider = Preconditions.checkNotNull(celTypeProvider); + this.legacyTypeProvider = legacyTypeProvider; + this.decls.add(Preconditions.checkNotNull(declGroup)); + } } diff --git a/checker/src/main/java/dev/cel/checker/ExprChecker.java b/checker/src/main/java/dev/cel/checker/ExprChecker.java index f72919f4e..c156e12e2 100644 --- a/checker/src/main/java/dev/cel/checker/ExprChecker.java +++ b/checker/src/main/java/dev/cel/checker/ExprChecker.java @@ -50,12 +50,16 @@ import dev.cel.common.types.CelKind; import dev.cel.common.types.CelProtoTypes; import dev.cel.common.types.CelType; +import dev.cel.common.types.CelTypeProvider; import dev.cel.common.types.CelTypes; +import dev.cel.common.types.EnumType; import dev.cel.common.types.ListType; import dev.cel.common.types.MapType; import dev.cel.common.types.OptionalType; import dev.cel.common.types.ProtoMessageType; import dev.cel.common.types.SimpleType; +import dev.cel.common.types.StructType; +import dev.cel.common.types.StructTypeReference; import dev.cel.common.types.TypeType; import java.util.ArrayList; import java.util.HashSet; @@ -168,7 +172,10 @@ public static CelAbstractSyntaxTree typecheck( } private final Env env; - private final TypeProvider typeProvider; + private final CelTypeProvider celTypeProvider; + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + private final @Nullable TypeProvider legacyTypeProvider; private final CelContainer container; private final Map positionMap; private final InferenceContext inferenceContext; @@ -177,25 +184,6 @@ public static CelAbstractSyntaxTree typecheck( private final boolean namespacedDeclarations; private final Set extensions; - private ExprChecker( - Env env, - CelContainer container, - Map positionMap, - InferenceContext inferenceContext, - boolean compileTimeOverloadResolution, - boolean homogeneousLiterals, - boolean namespacedDeclarations) { - this.env = checkNotNull(env); - this.typeProvider = env.getTypeProvider(); - this.positionMap = checkNotNull(positionMap); - this.container = checkNotNull(container); - this.inferenceContext = checkNotNull(inferenceContext); - this.compileTimeOverloadResolution = compileTimeOverloadResolution; - this.homogeneousLiterals = homogeneousLiterals; - this.namespacedDeclarations = namespacedDeclarations; - this.extensions = new HashSet<>(); - } - /** Visit the {@code expr} value, routing to overloads based on the kind of expression. */ public void visit(CelMutableExpr expr) { switch (expr.getKind()) { @@ -416,7 +404,7 @@ private void visit(CelMutableExpr expr, CelMutableStruct struct) { visit(value); CelType fieldType = - getFieldType(entry.id(), getPosition(entry), messageType, entry.fieldKey()).celType(); + getFieldType(entry.id(), getPosition(entry), messageType, entry.fieldKey()); CelType valueType = env.getType(value); if (entry.optionalEntry()) { if (valueType instanceof OptionalType) { @@ -691,14 +679,15 @@ private CelType visitSelectField( if (!Types.isDynOrError(operandType)) { if (operandType.kind().equals(CelKind.STRUCT)) { - TypeProvider.FieldType fieldType = - getFieldType(expr.id(), getPosition(expr), operandType, field); - ProtoMessageType protoMessageType = resolveProtoMessageType(operandType); - if (protoMessageType != null && protoMessageType.isJsonName(field)) { - extensions.add(JSON_NAME_EXTENSION); + CelType fieldType = getFieldType(expr.id(), getPosition(expr), operandType, field); + if (!fieldType.equals(SimpleType.ERROR)) { + ProtoMessageType protoMessageType = resolveProtoMessageType(operandType); + if (protoMessageType != null && protoMessageType.isJsonName(field)) { + extensions.add(JSON_NAME_EXTENSION); + } } // Type of the field - resultType = fieldType.celType(); + resultType = fieldType; } else if (operandType.kind().equals(CelKind.MAP)) { resultType = ((MapType) operandType).valueType(); } else if (operandType.kind().equals(CelKind.TYPE_PARAM)) { @@ -738,20 +727,10 @@ private CelType visitSelectField( if (operandType.kind().equals(CelKind.STRUCT)) { // This is either a StructTypeReference or just a Struct. Attempt to search for - // ProtoMessageType that may exist in in the type provider. - TypeType typeDef = - typeProvider - .lookupCelType(operandType.name()) - .filter(t -> t instanceof TypeType) - .map(TypeType.class::cast) - .orElse(null); - if (typeDef == null || typeDef.parameters().size() != 1) { - return null; - } - - CelType maybeProtoMessageType = typeDef.parameters().get(0); - if (maybeProtoMessageType instanceof ProtoMessageType) { - return (ProtoMessageType) maybeProtoMessageType; + // ProtoMessageType that may exist in the type provider. + CelType resolvedType = celTypeProvider.findType(operandType.name()).orElse(null); + if (resolvedType instanceof ProtoMessageType) { + return (ProtoMessageType) resolvedType; } } @@ -795,34 +774,106 @@ private void visitOptionalCall(CelMutableExpr expr, CelMutableCall call) { } } - /** Returns the field type give a type instance and field name. */ - private TypeProvider.FieldType getFieldType( - long exprId, int position, CelType type, String fieldName) { + /** Returns the field type given a type instance and field name. */ + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + private CelType getFieldType(long exprId, int position, CelType type, String fieldName) { String typeName = type.name(); - if (typeProvider.lookupCelType(typeName).isPresent()) { - TypeProvider.FieldType fieldType = typeProvider.lookupFieldType(type, fieldName); - if (fieldType != null) { - return fieldType; + CelType resolvedType = celTypeProvider.findType(typeName).orElse(null); + if (resolvedType instanceof StructType) { + StructType structType = (StructType) resolvedType; + StructType.Field field = structType.findField(fieldName).orElse(null); + if (field != null) { + return normalizeFieldType(field.type()); + } + if (structType instanceof ProtoMessageType) { + ProtoMessageType.Extension extension = + ((ProtoMessageType) structType).findExtension(fieldName).orElse(null); + if (extension != null) { + return normalizeFieldType(extension.type()); + } } - TypeProvider.ExtensionFieldType extensionFieldType = - typeProvider.lookupExtensionType(fieldName); - if (extensionFieldType != null) { - return extensionFieldType.fieldType(); + if (legacyTypeProvider != null) { + TypeProvider.ExtensionFieldType extensionFieldType = + legacyTypeProvider.lookupExtensionType(fieldName); + if (extensionFieldType != null + && extensionFieldType.messageType().getMessageType().equals(typeName)) { + return extensionFieldType.fieldType().celType(); + } } env.reportError(exprId, position, "undefined field '%s'", fieldName); - } else { - // Proto message was added as a variable to the environment but the descriptor was not - // provided - String errorMessage = - String.format("Message type resolution failure while referencing field '%s'.", fieldName); - if (type.kind().equals(CelKind.STRUCT)) { - errorMessage += - String.format( - " Ensure that the descriptor for type '%s' was added to the environment", typeName); + return SimpleType.ERROR; + } + + if (legacyTypeProvider != null && legacyTypeProvider.lookupCelType(typeName).isPresent()) { + Optional legacyFieldType = lookupLegacyFieldType(type, fieldName); + if (legacyFieldType.isPresent()) { + return legacyFieldType.get(); } - env.reportError(exprId, position, errorMessage, fieldName, typeName); + env.reportError(exprId, position, "undefined field '%s'", fieldName); + return SimpleType.ERROR; + } + + // Message/Struct was added as a variable to the environment but the descriptor was not + // provided + String errorMessage = + String.format("Message type resolution failure while referencing field '%s'.", fieldName); + if (type.kind().equals(CelKind.STRUCT)) { + errorMessage += + String.format( + " Ensure that the descriptor for type '%s' was added to the environment", typeName); } - return ERROR; + env.reportError(exprId, position, errorMessage); + return SimpleType.ERROR; + } + + private static CelType normalizeFieldType(CelType celType) { + if (celType instanceof EnumType) { + return SimpleType.INT; + } + if (celType instanceof StructType) { + return StructTypeReference.create(celType.name()); + } + if (celType instanceof ListType) { + ListType listType = (ListType) celType; + if (listType.hasElemType()) { + CelType normalizedElemType = normalizeFieldType(listType.elemType()); + if (!normalizedElemType.equals(listType.elemType())) { + return ListType.create(normalizedElemType); + } + } + return listType; + } + if (celType instanceof MapType) { + MapType mapType = (MapType) celType; + CelType normalizedKeyType = normalizeFieldType(mapType.keyType()); + CelType normalizedValueType = normalizeFieldType(mapType.valueType()); + if (!normalizedKeyType.equals(mapType.keyType()) + || !normalizedValueType.equals(mapType.valueType())) { + return MapType.create(normalizedKeyType, normalizedValueType); + } + return mapType; + } + return celType; + } + + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + private Optional lookupLegacyFieldType(CelType type, String fieldName) { + if (legacyTypeProvider == null) { + return Optional.absent(); + } + TypeProvider.FieldType legacyFieldType = legacyTypeProvider.lookupFieldType(type, fieldName); + if (legacyFieldType != null) { + return Optional.of(legacyFieldType.celType()); + } + TypeProvider.ExtensionFieldType extensionFieldType = + legacyTypeProvider.lookupExtensionType(fieldName); + if (extensionFieldType != null + && extensionFieldType.messageType().getMessageType().equals(type.name())) { + return Optional.of(extensionFieldType.fieldType().celType()); + } + return Optional.absent(); } /** Checks compatibility of joined types, and returns the most general common type. */ @@ -867,6 +918,28 @@ private int getPosition(CelMutableStruct.Entry entry) { return pos == null ? 0 : pos; } + // TypeProvider is deprecated, but preserved for backwards compatibility. + @SuppressWarnings("deprecation") + private ExprChecker( + Env env, + CelContainer container, + Map positionMap, + InferenceContext inferenceContext, + boolean compileTimeOverloadResolution, + boolean homogeneousLiterals, + boolean namespacedDeclarations) { + this.env = checkNotNull(env); + this.celTypeProvider = env.getCelTypeProvider(); + this.legacyTypeProvider = env.getTypeProvider(); + this.positionMap = checkNotNull(positionMap); + this.container = checkNotNull(container); + this.inferenceContext = checkNotNull(inferenceContext); + this.compileTimeOverloadResolution = compileTimeOverloadResolution; + this.homogeneousLiterals = homogeneousLiterals; + this.namespacedDeclarations = namespacedDeclarations; + this.extensions = new HashSet<>(); + } + /** Helper object for holding an overload resolution result. */ @AutoValue protected abstract static class OverloadResolution { @@ -882,7 +955,4 @@ public static OverloadResolution of(CelReference reference, CelType type) { return new AutoValue_ExprChecker_OverloadResolution(reference, type); } } - - /** Helper object to represent a {@link TypeProvider.FieldType} lookup failure. */ - private static final TypeProvider.FieldType ERROR = TypeProvider.FieldType.of(Types.ERROR); } diff --git a/checker/src/main/java/dev/cel/checker/TypeFormatter.java b/checker/src/main/java/dev/cel/checker/TypeFormatter.java index 3cdd1a511..19db6b34d 100644 --- a/checker/src/main/java/dev/cel/checker/TypeFormatter.java +++ b/checker/src/main/java/dev/cel/checker/TypeFormatter.java @@ -14,19 +14,18 @@ package dev.cel.checker; -import dev.cel.expr.Type; import dev.cel.common.annotations.Internal; import dev.cel.common.types.CelType; import dev.cel.common.types.CelTypes; import org.jspecify.annotations.Nullable; /** - * Class to format {@link Type} objects into {@code String} values. + * Class to format {@link CelType} objects into {@code String} values. * *

CEL Library Internals. Do Not Use. */ @Internal -public final class TypeFormatter { +final class TypeFormatter { /** Format a function string from the {@code argTypes} and {@code isInstance} information */ static String formatFunction(Iterable argTypes, boolean isInstance) { diff --git a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java index 92a70c2d6..0ec3d98bf 100644 --- a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java +++ b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java @@ -15,25 +15,35 @@ package dev.cel.checker; import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertThrows; import com.google.common.collect.ImmutableList; +import com.google.protobuf.Duration; +import com.google.protobuf.FieldMask; +import com.google.testing.junit.testparameterinjector.TestParameter; +import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.checker.CelStandardDeclarations.StandardFunction; +import dev.cel.common.CelAbstractSyntaxTree; import dev.cel.common.CelContainer; import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelOptions; import dev.cel.common.CelOverloadDecl; +import dev.cel.common.CelValidationException; import dev.cel.common.CelVarDecl; import dev.cel.common.types.CelType; import dev.cel.common.types.CelTypeProvider; +import dev.cel.common.types.ListType; +import dev.cel.common.types.MapType; import dev.cel.common.types.SimpleType; +import dev.cel.common.types.StructTypeReference; +import dev.cel.compiler.CelCompiler; import dev.cel.compiler.CelCompilerFactory; import dev.cel.expr.conformance.proto3.TestAllTypes; import java.util.Optional; import org.junit.Test; import org.junit.runner.RunWith; -import org.junit.runners.JUnit4; -@RunWith(JUnit4.class) +@RunWith(TestParameterInjector.class) public class CelCheckerLegacyImplTest { @Test @@ -148,4 +158,106 @@ public void toCheckerBuilder_collectionProperties_areImmutable() { assertThat(newCheckerBuilder.fileTypes().build()).isEmpty(); assertThat(newCheckerBuilder.checkerLibraries().build()).isEmpty(); } + + @Test + // TypeProvider is deprecated, but tested for backwards compatibility. + @SuppressWarnings("deprecation") + public void check_wellKnownTypeStructCreation_withLegacyTypeProvider_success() throws Exception { + TypeProvider legacyTypeProvider = + new DescriptorTypeProvider(ImmutableList.of(Duration.getDescriptor())); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder().setTypeProvider(legacyTypeProvider).build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Duration{seconds: 10, nanos: 20}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); + } + + @Test + // TypeProvider is deprecated, but tested for backwards compatibility. + @SuppressWarnings("deprecation") + public void check_protoTypeMask_failsClosedWithLegacyTypeProvider() throws Exception { + TypeProvider legacyTypeProvider = + new DescriptorTypeProvider(ImmutableList.of(TestAllTypes.getDescriptor())); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(TestAllTypes.getDescriptor()) + .addProtoTypeMasks( + ProtoTypeMask.of( + "cel.expr.conformance.proto3.TestAllTypes", + FieldMask.newBuilder().addPaths("single_int32").build())) + .setTypeProvider(legacyTypeProvider) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> + celCompiler + .compile( + "cel.expr.conformance.proto3.TestAllTypes{single_int32: 1, single_int64:" + + " 2}") + .getAst()); + + assertThat(e).hasMessageThat().contains("undefined field 'single_int64'"); + } + + @Test + public void lookupEnumValue_modernTypeProvider_success( + @TestParameter({ + "cel.expr.conformance.proto3.TestAllTypes.NestedEnum.BAZ == 2", + ".cel.expr.conformance.proto3.TestAllTypes.NestedEnum.BAZ == 2" + }) + String expr) + throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(TestAllTypes.getDescriptor()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile(expr).getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.BOOL); + } + + private enum FieldTypeTestCase { + SINGLE_ENUM("msg.single_nested_enum", SimpleType.INT), + REPEATED_ENUM("msg.repeated_nested_enum", ListType.create(SimpleType.INT)), + MAP_ENUM("msg.map_bool_enum", MapType.create(SimpleType.BOOL, SimpleType.INT)), + SINGLE_MESSAGE( + "msg.single_nested_message", + StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes.NestedMessage")), + REPEATED_MESSAGE( + "msg.repeated_nested_message", + ListType.create( + StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"))), + MAP_MESSAGE( + "msg.map_bool_message", + MapType.create( + SimpleType.BOOL, + StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"))); + + private final String expr; + private final CelType expectedType; + + FieldTypeTestCase(String expr, CelType expectedType) { + this.expr = expr; + this.expectedType = expectedType; + } + } + + @Test + public void check_enumAndMessageFields_retainsIntAndStructTypeReference( + @TestParameter FieldTypeTestCase testCase) throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(TestAllTypes.getDescriptor()) + .addVar("msg", StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes")) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile(testCase.expr).getAst(); + + assertThat(ast.getResultType()).isEqualTo(testCase.expectedType); + } } diff --git a/common/src/main/java/dev/cel/common/types/CelTypeProvider.java b/common/src/main/java/dev/cel/common/types/CelTypeProvider.java index c452a9dee..f662bffa6 100644 --- a/common/src/main/java/dev/cel/common/types/CelTypeProvider.java +++ b/common/src/main/java/dev/cel/common/types/CelTypeProvider.java @@ -47,6 +47,7 @@ public interface CelTypeProvider { @Immutable final class CombinedCelTypeProvider implements CelTypeProvider { + private final ImmutableList typeProviders; private final ImmutableMap allTypes; public CombinedCelTypeProvider(CelTypeProvider first, CelTypeProvider second) { @@ -54,6 +55,7 @@ public CombinedCelTypeProvider(CelTypeProvider first, CelTypeProvider second) { } public CombinedCelTypeProvider(ImmutableList typeProviders) { + this.typeProviders = ImmutableList.copyOf(typeProviders); Map allTypes = new LinkedHashMap<>(); typeProviders.forEach( typeProvider -> @@ -68,7 +70,17 @@ public ImmutableCollection types() { @Override public Optional findType(String typeName) { - return Optional.ofNullable(allTypes.get(typeName)); + CelType type = allTypes.get(typeName); + if (type != null) { + return Optional.of(type); + } + for (CelTypeProvider typeProvider : typeProviders) { + Optional resolvedType = typeProvider.findType(typeName); + if (resolvedType.isPresent()) { + return resolvedType; + } + } + return Optional.empty(); } } }