From e97bcaf5efd68450ee2557cfc8460b92289e05a4 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Fri, 2 Oct 2026 13:00:03 -0700 Subject: [PATCH] Fix type-checker to respect the ordering of type providers (protobuf then user provided) PiperOrigin-RevId: 992491032 --- .../dev/cel/checker/CelCheckerLegacyImpl.java | 2 +- .../cel/checker/CelCheckerLegacyImplTest.java | 43 +++++++++++++++++++ .../runtime/DescriptorTypeResolverTest.java | 3 -- 3 files changed, 44 insertions(+), 4 deletions(-) diff --git a/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java b/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java index cf11013c7..394389dfb 100644 --- a/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java +++ b/checker/src/main/java/dev/cel/checker/CelCheckerLegacyImpl.java @@ -446,7 +446,7 @@ public CelCheckerLegacyImpl build() { } else if (celTypeProvider != null) { messageTypeProvider = new CelTypeProvider.CombinedCelTypeProvider( - ImmutableList.of(celTypeProvider, messageTypeProvider)); + ImmutableList.of(messageTypeProvider, celTypeProvider)); } // Configure the declaration set, and possibly alter the type provider if ProtoDecl values diff --git a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java index 6771c1278..88a78e9ca 100644 --- a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java +++ b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java @@ -19,6 +19,7 @@ import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; import com.google.protobuf.Duration; import com.google.protobuf.FieldMask; import com.google.protobuf.Timestamp; @@ -40,6 +41,7 @@ import dev.cel.common.types.ListType; import dev.cel.common.types.MapType; 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.compiler.CelCompiler; @@ -409,6 +411,47 @@ public Optional findType(String typeName) { assertThat(ast.getResultType()).isEqualTo(preWrappedType); } + @Test + public void check_combinedTypeProviders_protoMessageTakesPrecedenceOverCustom() throws Exception { + StructType shadowingStructType = + StructType.create( + TestAllTypes.getDescriptor().getFullName(), + ImmutableSet.of("single_int64"), + fieldName -> Optional.of(SimpleType.STRING)); + StructType customOnlyStructType = + StructType.create( + "custom.CustomStruct", + ImmutableSet.of("custom_field"), + fieldName -> Optional.of(SimpleType.STRING)); + CelTypeProvider customTypeProvider = + new CelTypeProvider() { + @Override + public ImmutableList types() { + return ImmutableList.of(shadowingStructType, customOnlyStructType); + } + + @Override + public Optional findType(String typeName) { + return types().stream().filter(t -> t.name().equals(typeName)).findFirst(); + } + }; + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(TestAllTypes.getDescriptor()) + .setTypeProvider(customTypeProvider) + .build(); + + CelAbstractSyntaxTree protoAst = + celCompiler + .compile("cel.expr.conformance.proto3.TestAllTypes{single_int64: 1}.single_int64") + .getAst(); + CelAbstractSyntaxTree customAst = + celCompiler.compile("custom.CustomStruct{custom_field: 'hello'}.custom_field").getAst(); + + assertThat(protoAst.getResultType()).isEqualTo(SimpleType.INT); + assertThat(customAst.getResultType()).isEqualTo(SimpleType.STRING); + } + private enum FieldTypeTestCase { REPEATED_PRIMITIVE("msg.repeated_int64", ListType.create(SimpleType.INT)), MAP_PRIMITIVE("msg.map_string_string", MapType.create(SimpleType.STRING, SimpleType.STRING)), diff --git a/runtime/src/test/java/dev/cel/runtime/DescriptorTypeResolverTest.java b/runtime/src/test/java/dev/cel/runtime/DescriptorTypeResolverTest.java index 42f4dddc9..051186e74 100644 --- a/runtime/src/test/java/dev/cel/runtime/DescriptorTypeResolverTest.java +++ b/runtime/src/test/java/dev/cel/runtime/DescriptorTypeResolverTest.java @@ -45,9 +45,6 @@ public class DescriptorTypeResolverTest { private static final Cel CEL = CelFactory.plannerCelBuilder() .setTypeProvider(PROTO_MESSAGE_TYPE_PROVIDER) - // TODO: Replace setValueProvider with - // addMessageTypes(TestAllTypes.getDescriptor()) once CelRuntimeImpl prioritizes custom - // CelTypeProvider over its internal messageTypeProvider. .setValueProvider( (structType, fields) -> structType.equals(TestAllTypes.getDescriptor().getFullName())