From bb6a118391fec95473d7a0b94f9181c19700205e Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Thu, 17 Sep 2026 13:20:54 -0700 Subject: [PATCH] Modernize Env and ExprChecker to use CelTypeProvider PiperOrigin-RevId: 983383517 --- checker/BUILD.bazel | 11 +- .../src/main/java/dev/cel/checker/BUILD.bazel | 33 +- .../dev/cel/checker/CelCheckerLegacyImpl.java | 32 +- .../src/main/java/dev/cel/checker/Env.java | 229 +++++++++----- .../java/dev/cel/checker/ExprChecker.java | 200 ++++++++---- .../java/dev/cel/checker/TypeFormatter.java | 5 +- .../cel/checker/CelCheckerLegacyImplTest.java | 291 +++++++++++++++++- .../dev/cel/common/types/CelTypeProvider.java | 29 +- .../types/ProtoMessageTypeProviderTest.java | 39 +++ 9 files changed, 689 insertions(+), 180 deletions(-) diff --git a/checker/BUILD.bazel b/checker/BUILD.bazel index ac00ddff1..306aa9a4d 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"], @@ -35,7 +36,10 @@ java_library( java_library( name = "checker_legacy_environment", deprecation = "See go/cel-java-migration-guide. Please use CEL-Java Fluent APIs //compiler instead", - exports = ["//checker/src/main/java/dev/cel/checker:checker_legacy_environment"], + exports = [ + ":type_provider_legacy", + "//checker/src/main/java/dev/cel/checker:checker_legacy_environment", + ], ) java_library( @@ -53,3 +57,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..1050de508 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 = [ @@ -25,14 +26,11 @@ CHECKER_BUILDER_SOURCES = [ # keep sorted CHECKER_LEGACY_ENV_SOURCES = [ - "DescriptorTypeProvider.java", "Env.java", "ExprChecker.java", "ExprVisitor.java", "InferenceContext.java", "TypeFormatter.java", - "TypeProvider.java", - "Types.java", ] java_library( @@ -69,7 +67,7 @@ java_library( ":checker_legacy_environment", ":proto_type_mask", ":standard_decl", - ":type_provider_legacy_impl", + ":type_provider_legacy", "//common:cel_ast", "//common:cel_descriptor_util", "//common:cel_function_decl", @@ -102,9 +100,9 @@ java_library( tags = [ ], deps = [ - ":checker_legacy_environment", ":proto_type_mask", ":standard_decl", + ":type_provider_legacy", "//common:cel_ast", "//common:cel_function_decl", "//common:cel_validation_result", @@ -137,7 +135,7 @@ java_library( tags = [ ], deps = [ - ":checker_legacy_environment", + ":type_provider_legacy", "//common/annotations", "//common/types", "//common/types:cel_proto_types", @@ -156,6 +154,7 @@ java_library( ], deps = [ ":standard_decl", + ":type_provider_legacy", "//:auto_value", "//common:cel_ast", "//common:cel_function_decl", @@ -173,7 +172,6 @@ java_library( "//common/ast:expr_converter", "//common/ast:mutable_expr", "//common/internal:errors", - "//common/internal:file_descriptor_converter", "//common/types", "//common/types:cel_proto_types", "//common/types:cel_types", @@ -183,7 +181,6 @@ java_library( "@cel_spec//proto/cel/expr:syntax_java_proto", "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", - "@maven//:com_google_protobuf_protobuf_java", "@maven//:org_jspecify_jspecify", ], ) @@ -235,3 +232,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..cf11013c7 100644 --- a/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java +++ b/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java @@ -74,7 +74,7 @@ public final class CelCheckerLegacyImpl implements CelChecker, EnvVisitable { private final Optional expectedResultType; @SuppressWarnings("Immutable") - private final TypeProvider typeProvider; + private final @Nullable TypeProvider typeProvider; private final CelTypeProvider celTypeProvider; private final boolean standardEnvironmentEnabled; @@ -124,6 +124,10 @@ public CelCheckerBuilder toCheckerBuilder() { .addFileTypes(fileDescriptors) .addProtoTypeMasks(protoTypeMasks); + if (typeProvider != null) { + builder.setTypeProvider(typeProvider); + } + if (expectedResultType.isPresent()) { builder.setResultType(expectedResultType.get()); } @@ -162,11 +166,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); @@ -190,7 +196,7 @@ public static final class Builder implements CelCheckerBuilder { private CelContainer container; private CelOptions celOptions; private CelType expectedResultType; - private TypeProvider customTypeProvider; + private @Nullable TypeProvider customTypeProvider; private CelTypeProvider celTypeProvider; private boolean standardEnvironmentEnabled; private CelStandardDeclarations standardDeclarations; @@ -400,6 +406,11 @@ CelStandardDeclarations standardDeclarations() { return this.standardDeclarations; } + @VisibleForTesting + @Nullable TypeProvider customTypeProvider() { + return this.customTypeProvider; + } + @VisibleForTesting CelTypeProvider celTypeProvider() { return this.celTypeProvider; @@ -459,20 +470,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, @@ -499,7 +503,7 @@ private CelCheckerLegacyImpl( ImmutableSet identDeclarations, ImmutableSet functionDeclarations, Optional expectedResultType, - TypeProvider typeProvider, + @Nullable TypeProvider typeProvider, CelTypeProvider celTypeProvider, boolean standardEnvironmentEnabled, @Nullable CelStandardDeclarations overriddenStandardDeclarations, diff --git a/checker/src/main/java/dev/cel/checker/Env.java b/checker/src/main/java/dev/cel/checker/Env.java index 81fec362c..97908dca3 100644 --- a/checker/src/main/java/dev/cel/checker/Env.java +++ b/checker/src/main/java/dev/cel/checker/Env.java @@ -46,8 +46,11 @@ 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.SimpleType; +import dev.cel.common.types.TypeType; import dev.cel.parser.CelStandardMacro; import java.util.ArrayList; import java.util.HashMap; @@ -80,8 +83,13 @@ 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; + + private final @Nullable TypeProvider legacyTypeProvider; /** * Stack of declaration groups where each entry in stack represents a scope capable of hinding @@ -107,14 +115,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 +133,14 @@ static Env unconfigured(Errors errors, CelOptions celOptions) { return unconfigured(errors, new DescriptorTypeProvider(), celOptions); } + 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 +150,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 +171,16 @@ public static Env standard(Errors errors, TypeProvider typeProvider) { return standard(errors, typeProvider, LEGACY_TYPE_CHECKER_OPTIONS); } + 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 +194,17 @@ 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(celStandardDeclaration, errors, typeProvider, celOptions); + return standard( + newStandardDeclarations(celOptions), errors, EMPTY_TYPE_PROVIDER, typeProvider, celOptions); } - public static Env standard( + 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 +213,68 @@ 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}. */ + CelTypeProvider getCelTypeProvider() { + return celTypeProvider; + } + + /** + * Returns the {@code TypeProvider}, or {@code null} if only modern {@link CelTypeProvider} was + * configured. + * + * @deprecated Use {@link #getCelTypeProvider()} instead. + */ + @Deprecated + @Nullable TypeProvider getTypeProvider() { + return legacyTypeProvider; } /** @@ -478,7 +519,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 +535,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 +550,38 @@ 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; + } + } + return Optional.ofNullable(legacyTypeProvider).map(value -> value.lookupEnumValue(enumName)); + } + + private static CelType wrapAsTypeIfNeeded(CelType type) { + if (type instanceof TypeType) { + return type; + } + return TypeType.create(type); + } + /** * Lookup a local identifier by name. This searches only comprehension scopes, bypassing standard * environment or user-defined environment. @@ -889,37 +969,21 @@ public Decl build() { * *

Identifiers and functions can share the same declaration name, so a simple map will not * suffice for tracking declaration overloads. - * - *

Whether a given {@code DeclGroup} is mutable or immutable depends on whether the maps - * supplied as input to the group are standard {@code Map} implementations or {@code ImmutableMap} - * implementations. The {DeclGroup#immutableCopy} method is provided as a convenience to make it - * easy to create an instance of the group which will honor the developer's intent. */ public static class DeclGroup { private final Map idents; private final Map functions; - /** Construct an empty {@code DeclGroup}. */ - public DeclGroup() { - this(new HashMap<>(), new HashMap<>()); - } - - /** Construct a new {@code DeclGroup} from the input {@code idents} and {@code functions}. */ - public DeclGroup(Map idents, Map functions) { - this.functions = functions; - this.idents = idents; - } - /** * Get an immutable map of the identifiers in the {@code DeclGroup} keyed by declaration name. */ - public Map getIdents() { + public ImmutableMap getIdents() { return ImmutableMap.copyOf(idents); } /** Get an immutable map of the functions in the {@code DeclGroup} keyed by declaration name. */ - public Map getFunctions() { + public ImmutableMap getFunctions() { return ImmutableMap.copyOf(functions); } @@ -943,9 +1007,9 @@ public void putFunction(CelFunctionDecl function) { functions.put(function.name(), function); } - /** Create a copy of the {@code DeclGroup} with immutable identifier and function maps. */ - public DeclGroup immutableCopy() { - return new DeclGroup(getIdents(), getFunctions()); + private DeclGroup() { + this.idents = new HashMap<>(); + this.functions = new HashMap<>(); } } @@ -1015,4 +1079,17 @@ static CelType getWellKnownType(CelType type) { Preconditions.checkArgument(type.kind() == CelKind.STRUCT); return CelTypes.getWellKnownCelType(type.name()).get(); } + + 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..8a842ce7a 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; + + 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. */ + 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) { + Optional extensionType = + lookupLegacyExtensionType(legacyTypeProvider, typeName, fieldName); + if (extensionType.isPresent()) { + return extensionType.get(); + } } 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(legacyTypeProvider, type, fieldName); + if (legacyFieldType.isPresent()) { + return legacyFieldType.get(); + } + 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); + } + 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); + } } - env.reportError(exprId, position, errorMessage, fieldName, typeName); + 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 ERROR; + return 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); + if (legacyFieldType != null) { + return Optional.of(legacyFieldType.celType()); + } + return lookupLegacyExtensionType(legacyTypeProvider, type.name(), fieldName); + } + + private static Optional lookupLegacyExtensionType( + TypeProvider legacyTypeProvider, String typeName, String fieldName) { + TypeProvider.ExtensionFieldType extensionFieldType = + legacyTypeProvider.lookupExtensionType(fieldName); + if (extensionFieldType != null + && extensionFieldType.messageType().getMessageType().equals(typeName)) { + return Optional.of(extensionFieldType.fieldType().celType()); + } + return Optional.absent(); } /** Checks compatibility of joined types, and returns the most general common type. */ @@ -867,6 +918,26 @@ private int getPosition(CelMutableStruct.Entry entry) { return pos == null ? 0 : pos; } + 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 +953,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..3e557d0eb 100644 --- a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java +++ b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java @@ -15,25 +15,42 @@ 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.common.collect.ImmutableMap; +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.CelSource; +import dev.cel.common.CelValidationException; import dev.cel.common.CelVarDecl; +import dev.cel.common.ast.CelExpr; import dev.cel.common.types.CelType; import dev.cel.common.types.CelTypeProvider; +import dev.cel.common.types.EnumType; +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.common.types.TypeType; +import dev.cel.compiler.CelCompiler; import dev.cel.compiler.CelCompilerFactory; +import dev.cel.expr.conformance.proto2.TestAllTypesExtensions; +import dev.cel.expr.conformance.proto2.TestAllTypesProto; 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 @@ -66,6 +83,7 @@ public void toCheckerBuilder_singularFields_copied() { CelOptions celOptions = CelOptions.current().build(); CelContainer celContainer = CelContainer.ofName("foo"); CelType expectedResultType = SimpleType.BOOL; + TypeProvider legacyTypeProvider = new DescriptorTypeProvider(); CelTypeProvider customTypeProvider = new CelTypeProvider() { @Override @@ -83,6 +101,7 @@ public Optional findType(String typeName) { .setOptions(celOptions) .setContainer(celContainer) .setResultType(expectedResultType) + .setTypeProvider(legacyTypeProvider) .setTypeProvider(customTypeProvider) .setStandardEnvironmentEnabled(false) .setStandardDeclarations(subsetDecls); @@ -94,6 +113,7 @@ public Optional findType(String typeName) { assertThat(newCheckerBuilder.standardDeclarations()).isEqualTo(subsetDecls); assertThat(newCheckerBuilder.options()).isEqualTo(celOptions); assertThat(newCheckerBuilder.container()).isEqualTo(celContainer); + assertThat(newCheckerBuilder.customTypeProvider()).isEqualTo(legacyTypeProvider); assertThat(newCheckerBuilder.celTypeProvider()).isEqualTo(customTypeProvider); } @@ -148,4 +168,271 @@ public void toCheckerBuilder_collectionProperties_areImmutable() { assertThat(newCheckerBuilder.fileTypes().build()).isEmpty(); assertThat(newCheckerBuilder.checkerLibraries().build()).isEmpty(); } + + @Test + 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 + 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 check_extensionField_withModernTypeProvider_success() throws Exception { + CelChecker celChecker = + CelCompilerFactory.standardCelCheckerBuilder() + .addFileTypes(TestAllTypesProto.getDescriptor(), TestAllTypesExtensions.getDescriptor()) + .addVarDeclarations( + CelVarDecl.newVarDeclaration( + "msg", StructTypeReference.create("cel.expr.conformance.proto2.TestAllTypes"))) + .build(); + CelAbstractSyntaxTree parsedAst = + CelAbstractSyntaxTree.newParsedAst( + CelExpr.ofSelect( + 2L, + CelExpr.ofIdent(1L, "msg"), + "cel.expr.conformance.proto2.int32_ext", + /* isTestOnly= */ false), + CelSource.newBuilder().build()); + + CelAbstractSyntaxTree checkedAst = celChecker.check(parsedAst).getAst(); + + assertThat(checkedAst.getResultType()).isEqualTo(SimpleType.INT); + } + + @Test + public void check_extensionField_withModernMessageAndLegacyExtensionProvider_success() + throws Exception { + TypeProvider legacyExtensionProvider = + new DescriptorTypeProvider( + ImmutableList.of( + TestAllTypesProto.getDescriptor(), TestAllTypesExtensions.getDescriptor())); + CelChecker celChecker = + CelCompilerFactory.standardCelCheckerBuilder() + .addMessageTypes(TestAllTypesExtensions.int32Ext.getDescriptor().getContainingType()) + .setTypeProvider(legacyExtensionProvider) + .addVarDeclarations( + CelVarDecl.newVarDeclaration( + "msg", StructTypeReference.create("cel.expr.conformance.proto2.TestAllTypes"))) + .build(); + CelAbstractSyntaxTree parsedAst = + CelAbstractSyntaxTree.newParsedAst( + CelExpr.ofSelect( + 2L, + CelExpr.ofIdent(1L, "msg"), + "cel.expr.conformance.proto2.int32_ext", + /* isTestOnly= */ false), + CelSource.newBuilder().build()); + + CelAbstractSyntaxTree checkedAst = celChecker.check(parsedAst).getAst(); + + assertThat(checkedAst.getResultType()).isEqualTo(SimpleType.INT); + } + + @Test + public void check_fieldAndExtensionSelection_withLegacyTypeProviderOnly_success( + @TestParameter({"single_int64", "cel.expr.conformance.proto2.int32_ext"}) String fieldName) + throws Exception { + TypeProvider legacyTypeProvider = + new DescriptorTypeProvider( + ImmutableList.of( + TestAllTypesProto.getDescriptor(), TestAllTypesExtensions.getDescriptor())); + CelChecker celChecker = + CelCompilerFactory.standardCelCheckerBuilder() + .setTypeProvider(legacyTypeProvider) + .addVarDeclarations( + CelVarDecl.newVarDeclaration( + "msg", StructTypeReference.create("cel.expr.conformance.proto2.TestAllTypes"))) + .build(); + CelAbstractSyntaxTree parsedAst = + CelAbstractSyntaxTree.newParsedAst( + CelExpr.ofSelect(2L, CelExpr.ofIdent(1L, "msg"), fieldName, /* isTestOnly= */ false), + CelSource.newBuilder().build()); + + CelAbstractSyntaxTree checkedAst = celChecker.check(parsedAst).getAst(); + + assertThat(checkedAst.getResultType()).isEqualTo(SimpleType.INT); + } + + @Test + public void check_undefinedFieldSelection_withLegacyTypeProviderOnly_throws() { + TypeProvider legacyTypeProvider = + new DescriptorTypeProvider(ImmutableList.of(TestAllTypes.getDescriptor())); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setTypeProvider(legacyTypeProvider) + .addVar("msg", StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes")) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("msg.undefined_field").getAst()); + + assertThat(e).hasMessageThat().contains("undefined field 'undefined_field'"); + } + + @Test + public void check_undeclaredMessageTypeFieldSelection_throws() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addVar("msg", StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes")) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, () -> celCompiler.compile("msg.single_int64").getAst()); + + assertThat(e) + .hasMessageThat() + .contains( + "Message type resolution failure while referencing field 'single_int64'. Ensure that" + + " the descriptor for type 'cel.expr.conformance.proto3.TestAllTypes' was added to" + + " the environment"); + } + + @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); + } + + @Test + public void lookupEnumValue_dynamicTypeProviderKeyedByValueName_success() throws Exception { + CelTypeProvider dynamicEnumProvider = + new CelTypeProvider() { + @Override + public ImmutableList types() { + return ImmutableList.of(); + } + + @Override + public Optional findType(String typeName) { + if (typeName.equals("custom.Enum.VALUE")) { + return Optional.of(EnumType.create("custom.Enum", ImmutableMap.of("VALUE", 42))); + } + return Optional.empty(); + } + }; + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(TestAllTypes.getDescriptor()) + .setTypeProvider(dynamicEnumProvider) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("custom.Enum.VALUE == 42").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.BOOL); + } + + @Test + public void check_customTypeProviderReturningPreWrappedType_doesNotDoubleWrap() throws Exception { + TypeType preWrappedType = TypeType.create(StructTypeReference.create("custom.MyType")); + CelTypeProvider customTypeProvider = + new CelTypeProvider() { + @Override + public ImmutableList types() { + return ImmutableList.of(); + } + + @Override + public Optional findType(String typeName) { + return typeName.equals("custom.MyType") + ? Optional.of(preWrappedType) + : Optional.empty(); + } + }; + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder().setTypeProvider(customTypeProvider).build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("custom.MyType").getAst(); + + assertThat(ast.getResultType()).isEqualTo(preWrappedType); + } + + private enum FieldTypeTestCase { + REPEATED_PRIMITIVE("msg.repeated_int64", ListType.create(SimpleType.INT)), + MAP_PRIMITIVE("msg.map_string_string", MapType.create(SimpleType.STRING, SimpleType.STRING)), + 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..1366ba09a 100644 --- a/common/src/main/java/dev/cel/common/types/CelTypeProvider.java +++ b/common/src/main/java/dev/cel/common/types/CelTypeProvider.java @@ -21,6 +21,7 @@ import com.google.errorprone.annotations.Immutable; import java.util.LinkedHashMap; import java.util.Map; +import java.util.Objects; import java.util.Optional; /** @@ -47,28 +48,36 @@ public interface CelTypeProvider { @Immutable final class CombinedCelTypeProvider implements CelTypeProvider { + private final ImmutableList typeProviders; private final ImmutableMap allTypes; + @Override + public ImmutableCollection types() { + return allTypes.values(); + } + + @Override + public Optional findType(String typeName) { + for (CelTypeProvider typeProvider : typeProviders) { + Optional resolvedType = typeProvider.findType(typeName); + if (resolvedType.isPresent()) { + return resolvedType; + } + } + return Optional.empty(); + } + public CombinedCelTypeProvider(CelTypeProvider first, CelTypeProvider second) { this(ImmutableList.of(first, second)); } public CombinedCelTypeProvider(ImmutableList typeProviders) { + this.typeProviders = Objects.requireNonNull(typeProviders); Map allTypes = new LinkedHashMap<>(); typeProviders.forEach( typeProvider -> typeProvider.types().forEach(type -> allTypes.putIfAbsent(type.name(), type))); this.allTypes = ImmutableMap.copyOf(allTypes); } - - @Override - public ImmutableCollection types() { - return allTypes.values(); - } - - @Override - public Optional findType(String typeName) { - return Optional.ofNullable(allTypes.get(typeName)); - } } } diff --git a/common/src/test/java/dev/cel/common/types/ProtoMessageTypeProviderTest.java b/common/src/test/java/dev/cel/common/types/ProtoMessageTypeProviderTest.java index c9f9d9e21..adfb47088 100644 --- a/common/src/test/java/dev/cel/common/types/ProtoMessageTypeProviderTest.java +++ b/common/src/test/java/dev/cel/common/types/ProtoMessageTypeProviderTest.java @@ -257,6 +257,45 @@ public void types_combinedDuplicateProviderIsSameAsFirst() { assertThat(combined.types()).hasSize(proto3Provider.types().size()); } + @Test + public void findType_combinedWithDynamicProvider_resolvesFromDelegateWithFirstPrecedence() { + CelType dynamicType = StructTypeReference.create("custom.DynamicType"); + CelType overriddenTestAllTypes = + StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes"); + CelTypeProvider dynamicProvider = + new CelTypeProvider() { + @Override + public ImmutableList types() { + return ImmutableList.of(); + } + + @Override + public Optional findType(String typeName) { + if (typeName.equals("custom.DynamicType")) { + return Optional.of(dynamicType); + } + if (typeName.equals("cel.expr.conformance.proto3.TestAllTypes")) { + return Optional.of(overriddenTestAllTypes); + } + return Optional.empty(); + } + }; + CombinedCelTypeProvider dynamicFirst = + new CombinedCelTypeProvider(dynamicProvider, proto3Provider); + CombinedCelTypeProvider staticFirst = + new CombinedCelTypeProvider(proto3Provider, dynamicProvider); + + assertThat(dynamicFirst.findType("cel.expr.conformance.proto3.TestAllTypes")) + .hasValue(overriddenTestAllTypes); + assertThat(dynamicFirst.findType("cel.expr.conformance.proto3.TestAllTypes.NestedMessage")) + .isPresent(); + assertThat(dynamicFirst.findType("custom.DynamicType")).hasValue(dynamicType); + assertThat(dynamicFirst.findType("custom.UndefinedType")).isEmpty(); + assertThat(staticFirst.findType("cel.expr.conformance.proto3.TestAllTypes").get()) + .isInstanceOf(ProtoMessageType.class); + assertThat(staticFirst.findType("custom.DynamicType")).hasValue(dynamicType); + } + @Test public void findField_withJsonNameOption() { ProtoMessageTypeProvider typeProvider =