From 782eae03f8012328d61398ff18696bbce2478d8c Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Wed, 7 Oct 2026 16:36:08 -0700 Subject: [PATCH] Implement OptimizedSelectable interface on full protobuf message PiperOrigin-RevId: 995434087 --- .../java/dev/cel/common/values/BUILD.bazel | 2 + .../cel/common/values/ProtoMessageValue.java | 76 +++-- .../common/values/ProtoMessageValueTest.java | 262 +++++++++++++++--- 3 files changed, 291 insertions(+), 49 deletions(-) diff --git a/common/src/main/java/dev/cel/common/values/BUILD.bazel b/common/src/main/java/dev/cel/common/values/BUILD.bazel index 785a2ef32..d50e54fc5 100644 --- a/common/src/main/java/dev/cel/common/values/BUILD.bazel +++ b/common/src/main/java/dev/cel/common/values/BUILD.bazel @@ -280,7 +280,9 @@ java_library( ], deps = [ ":base_proto_cel_value_converter", + ":optimized_selectable", ":preadapted_list", + ":select_field", ":values", "//:auto_value", "//common:options", diff --git a/common/src/main/java/dev/cel/common/values/ProtoMessageValue.java b/common/src/main/java/dev/cel/common/values/ProtoMessageValue.java index 627bd2c1d..35331cdfc 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoMessageValue.java +++ b/common/src/main/java/dev/cel/common/values/ProtoMessageValue.java @@ -14,8 +14,9 @@ package dev.cel.common.values; +import static com.google.common.base.Preconditions.checkNotNull; + import com.google.auto.value.AutoValue; -import com.google.common.base.Preconditions; import com.google.errorprone.annotations.Immutable; import com.google.protobuf.Descriptors.Descriptor; import com.google.protobuf.Descriptors.FieldDescriptor; @@ -28,7 +29,8 @@ /** ProtoMessageValue is a struct value with protobuf support. */ @AutoValue @Immutable -public abstract class ProtoMessageValue extends StructValue { +public abstract class ProtoMessageValue extends StructValue + implements OptimizedSelectable { @Override public abstract Message value(); @@ -60,18 +62,28 @@ public Optional find(String field) { FieldDescriptor fieldDescriptor = findField(celDescriptorPool(), value().getDescriptorForType(), field); - // Selecting a field on a protobuf message yields a default value even if the field is not - // declared. Therefore, we must exhaustively test whether they are actually declared. - if (fieldDescriptor.isRepeated()) { - if (value().getRepeatedFieldCount(fieldDescriptor) == 0) { - return Optional.empty(); - } - } else if (!value().hasField(fieldDescriptor)) { - return Optional.empty(); - } + return findFieldValue(fieldDescriptor); + } - return Optional.of( - protoCelValueConverter().fromProtoMessageFieldToCelValue(value(), fieldDescriptor)); + @Override + public Object selectByFieldNumber(SelectField field) { + FieldDescriptor fieldDescriptor = findFieldByNumber(value().getDescriptorForType(), field); + + return protoCelValueConverter().fromProtoMessageFieldToCelValue(value(), fieldDescriptor); + } + + @Override + public boolean hasFieldByNumber(SelectField field) { + FieldDescriptor fieldDescriptor = findFieldByNumber(value().getDescriptorForType(), field); + + return isFieldPresent(fieldDescriptor); + } + + @Override + public Optional findByFieldNumber(SelectField field) { + FieldDescriptor fieldDescriptor = findFieldByNumber(value().getDescriptorForType(), field); + + return findFieldValue(fieldDescriptor); } public static ProtoMessageValue create( @@ -79,9 +91,9 @@ public static ProtoMessageValue create( CelDescriptorPool celDescriptorPool, ProtoCelValueConverter protoCelValueConverter, boolean enableJsonFieldNames) { - Preconditions.checkNotNull(value); - Preconditions.checkNotNull(celDescriptorPool); - Preconditions.checkNotNull(protoCelValueConverter); + checkNotNull(value); + checkNotNull(celDescriptorPool); + checkNotNull(protoCelValueConverter); return new AutoValue_ProtoMessageValue( value, StructTypeReference.create(value.getDescriptorForType().getFullName()), @@ -90,6 +102,36 @@ public static ProtoMessageValue create( enableJsonFieldNames); } + private Optional findFieldValue(FieldDescriptor fieldDescriptor) { + if (!isFieldPresent(fieldDescriptor)) { + return Optional.empty(); + } + + return Optional.of( + protoCelValueConverter().fromProtoMessageFieldToCelValue(value(), fieldDescriptor)); + } + + private boolean isFieldPresent(FieldDescriptor fieldDescriptor) { + // Selecting a field on a protobuf message yields a default value even if the field is not + // declared. Therefore, we must exhaustively test whether they are actually declared. + if (fieldDescriptor.isRepeated()) { + return value().getRepeatedFieldCount(fieldDescriptor) > 0; + } + return value().hasField(fieldDescriptor); + } + + private static FieldDescriptor findFieldByNumber(Descriptor descriptor, SelectField field) { + FieldDescriptor fieldDescriptor = descriptor.findFieldByNumber(field.fieldNumber()); + if (fieldDescriptor != null) { + return fieldDescriptor; + } + + throw new IllegalArgumentException( + String.format( + "field '%s' (number %d) is not declared in message '%s'", + field.fieldName(), field.fieldNumber(), descriptor.getFullName())); + } + private FieldDescriptor findField( CelDescriptorPool celDescriptorPool, Descriptor descriptor, String fieldName) { if (enableJsonFieldNames()) { @@ -114,4 +156,6 @@ private FieldDescriptor findField( "field '%s' is not declared in message '%s'", fieldName, descriptor.getFullName()))); } + + ProtoMessageValue() {} } diff --git a/common/src/test/java/dev/cel/common/values/ProtoMessageValueTest.java b/common/src/test/java/dev/cel/common/values/ProtoMessageValueTest.java index b5c29129b..a718af143 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoMessageValueTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoMessageValueTest.java @@ -21,12 +21,17 @@ import com.google.common.collect.ImmutableMap; import com.google.common.primitives.UnsignedLong; import com.google.protobuf.Any; +import com.google.protobuf.BoolValue; import com.google.protobuf.ByteString; +import com.google.protobuf.BytesValue; +import com.google.protobuf.DoubleValue; import com.google.protobuf.DynamicMessage; import com.google.protobuf.FieldMask; import com.google.protobuf.FloatValue; import com.google.protobuf.Int32Value; import com.google.protobuf.Int64Value; +import com.google.protobuf.ListValue; +import com.google.protobuf.StringValue; import com.google.protobuf.Struct; import com.google.protobuf.Timestamp; import com.google.protobuf.UInt32Value; @@ -48,6 +53,7 @@ import dev.cel.expr.conformance.proto2.TestAllTypesExtensions; import java.time.Duration; import java.time.Instant; +import java.util.Optional; import org.junit.Test; import org.junit.runner.RunWith; @@ -215,39 +221,12 @@ private enum SelectFieldTestCase { @Test public void selectField_success(@TestParameter SelectFieldTestCase testCase) { - TestAllTypes testAllTypes = - TestAllTypes.newBuilder() - .setSingleBool(true) - .setSingleInt32(4) - .setSingleInt64(5L) - .setSingleUint32(1) - .setSingleUint64(UnsignedLong.MAX_VALUE.longValue()) - .setSingleFloat(1.5f) - .setSingleDouble(2.5d) - .setSingleString("test") - .setSingleBytes(ByteString.copyFrom(new byte[] {0x01})) - .setSingleAny( - Any.pack(DynamicMessage.newBuilder(com.google.protobuf.BoolValue.of(true)).build())) - .setSingleDuration(com.google.protobuf.Duration.newBuilder().setSeconds(100)) - .setSingleTimestamp(Timestamp.newBuilder().setSeconds(100)) - .setSingleInt32Wrapper(Int32Value.of(5)) - .setSingleInt64Wrapper(Int64Value.of(10L)) - .setSingleUint32Wrapper(UInt32Value.of(1)) - .setSingleUint64Wrapper(UInt64Value.of(UnsignedLong.MAX_VALUE.longValue())) - .setSingleStringWrapper(com.google.protobuf.StringValue.of("hello")) - .setSingleFloatWrapper(FloatValue.of(7.5f)) - .setSingleDoubleWrapper(com.google.protobuf.DoubleValue.of(8.5d)) - .setSingleBytesWrapper( - com.google.protobuf.BytesValue.of(ByteString.copyFrom(new byte[] {0x02}))) - .addRepeatedInt64(5L) - .putMapStringString("a", "b") - .setStandaloneMessage(NestedMessage.getDefaultInstance()) - .setStandaloneEnum(NestedEnum.BAR) - .build(); - ProtoMessageValue protoMessageValue = ProtoMessageValue.create( - testAllTypes, DefaultDescriptorPool.INSTANCE, PROTO_CEL_VALUE_CONVERTER, false); + createPopulatedTestAllTypes(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); assertThat(protoMessageValue.select(testCase.fieldName)).isEqualTo(testCase.value); } @@ -361,7 +340,7 @@ private enum SelectFieldJsonValueTestCase { LIST( Value.newBuilder() .setListValue( - com.google.protobuf.ListValue.newBuilder() + ListValue.newBuilder() .addValues(Value.newBuilder().setStringValue("test").build()) .build()) .build(), @@ -410,7 +389,7 @@ public void selectField_jsonList() { TestAllTypes testAllTypes = TestAllTypes.newBuilder() .setListValue( - com.google.protobuf.ListValue.newBuilder() + ListValue.newBuilder() .addValues(Value.newBuilder().setBoolValue(false).build()) .build()) .build(); @@ -457,4 +436,221 @@ public void findField_jsonName_success() { assertThat(protoMessageValue.find("singleInt32")).isPresent(); } + + @Test + public void selectByFieldNumber_success(@TestParameter SelectFieldTestCase testCase) { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + createPopulatedTestAllTypes(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + int fieldNumber = TestAllTypes.getDescriptor().findFieldByName(testCase.fieldName).getNumber(); + SelectField selectField = SelectField.create(fieldNumber, "renamed_" + testCase.fieldName); + + Object result = protoMessageValue.selectByFieldNumber(selectField); + + assertThat(result).isEqualTo(testCase.value); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum UnsetFieldDefaultTestCase { + PROTO2_CUSTOM_INT32(TestAllTypes.SINGLE_INT32_FIELD_NUMBER, "single_int32", -32L), + PROTO2_CUSTOM_STRING(TestAllTypes.SINGLE_STRING_FIELD_NUMBER, "single_string", "empty"), + WRAPPER_UNSET( + TestAllTypes.SINGLE_INT64_WRAPPER_FIELD_NUMBER, + "single_int64_wrapper", + NullValue.NULL_VALUE), + REPEATED_EMPTY(TestAllTypes.REPEATED_INT64_FIELD_NUMBER, "repeated_int64", ImmutableList.of()), + MAP_EMPTY(TestAllTypes.MAP_STRING_STRING_FIELD_NUMBER, "map_string_string", ImmutableMap.of()); + + private final int fieldNumber; + private final String fieldName; + private final Object expectedDefault; + + UnsetFieldDefaultTestCase(int fieldNumber, String fieldName, Object expectedDefault) { + this.fieldNumber = fieldNumber; + this.fieldName = fieldName; + this.expectedDefault = expectedDefault; + } + } + + @Test + public void selectByFieldNumber_unsetField_returnsDefault( + @TestParameter UnsetFieldDefaultTestCase testCase) { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + TestAllTypes.getDefaultInstance(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + SelectField selectField = + SelectField.create(testCase.fieldNumber, "renamed_" + testCase.fieldName); + + Object result = protoMessageValue.selectByFieldNumber(selectField); + + assertThat(result).isEqualTo(testCase.expectedDefault); + } + + @Test + public void selectByFieldNumber_undeclaredField_throwsException() { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + TestAllTypes.getDefaultInstance(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + SelectField undeclaredField = SelectField.create(99999L, "bogus"); + + IllegalArgumentException exception = + assertThrows( + IllegalArgumentException.class, + () -> protoMessageValue.selectByFieldNumber(undeclaredField)); + + assertThat(exception) + .hasMessageThat() + .isEqualTo( + "field 'bogus' (number 99999) is not declared in message" + + " 'cel.expr.conformance.proto2.TestAllTypes'"); + } + + @Test + public void hasFieldByNumber_fieldIsSet_returnsTrue(@TestParameter SelectFieldTestCase testCase) { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + createPopulatedTestAllTypes(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + int fieldNumber = TestAllTypes.getDescriptor().findFieldByName(testCase.fieldName).getNumber(); + SelectField selectField = SelectField.create(fieldNumber, "renamed_" + testCase.fieldName); + + boolean result = protoMessageValue.hasFieldByNumber(selectField); + + assertThat(result).isTrue(); + } + + @Test + public void hasFieldByNumber_fieldIsUnset_returnsFalse( + @TestParameter SelectFieldTestCase testCase) { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + TestAllTypes.getDefaultInstance(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + int fieldNumber = TestAllTypes.getDescriptor().findFieldByName(testCase.fieldName).getNumber(); + SelectField selectField = SelectField.create(fieldNumber, "renamed_" + testCase.fieldName); + + boolean result = protoMessageValue.hasFieldByNumber(selectField); + + assertThat(result).isFalse(); + } + + @Test + public void hasFieldByNumber_undeclaredField_throwsException() { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + TestAllTypes.getDefaultInstance(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + SelectField undeclaredField = SelectField.create(99999L, "bogus"); + + IllegalArgumentException exception = + assertThrows( + IllegalArgumentException.class, + () -> protoMessageValue.hasFieldByNumber(undeclaredField)); + + assertThat(exception) + .hasMessageThat() + .isEqualTo( + "field 'bogus' (number 99999) is not declared in message" + + " 'cel.expr.conformance.proto2.TestAllTypes'"); + } + + @Test + public void findByFieldNumber_fieldIsSet_returnsValue( + @TestParameter SelectFieldTestCase testCase) { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + createPopulatedTestAllTypes(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + int fieldNumber = TestAllTypes.getDescriptor().findFieldByName(testCase.fieldName).getNumber(); + SelectField selectField = SelectField.create(fieldNumber, "renamed_" + testCase.fieldName); + + Optional result = protoMessageValue.findByFieldNumber(selectField); + + assertThat(result).hasValue(testCase.value); + } + + @Test + public void findByFieldNumber_fieldIsUnset_returnsEmpty( + @TestParameter SelectFieldTestCase testCase) { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + TestAllTypes.getDefaultInstance(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + int fieldNumber = TestAllTypes.getDescriptor().findFieldByName(testCase.fieldName).getNumber(); + SelectField selectField = SelectField.create(fieldNumber, "renamed_" + testCase.fieldName); + + Optional result = protoMessageValue.findByFieldNumber(selectField); + + assertThat(result).isEmpty(); + } + + @Test + public void findByFieldNumber_undeclaredField_throwsException() { + ProtoMessageValue protoMessageValue = + ProtoMessageValue.create( + TestAllTypes.getDefaultInstance(), + DefaultDescriptorPool.INSTANCE, + PROTO_CEL_VALUE_CONVERTER, + /* enableJsonFieldNames= */ false); + SelectField undeclaredField = SelectField.create(99999L, "bogus"); + + IllegalArgumentException exception = + assertThrows( + IllegalArgumentException.class, + () -> protoMessageValue.findByFieldNumber(undeclaredField)); + + assertThat(exception) + .hasMessageThat() + .isEqualTo( + "field 'bogus' (number 99999) is not declared in message" + + " 'cel.expr.conformance.proto2.TestAllTypes'"); + } + + private static TestAllTypes createPopulatedTestAllTypes() { + return TestAllTypes.newBuilder() + .setSingleBool(true) + .setSingleInt32(4) + .setSingleInt64(5L) + .setSingleUint32(1) + .setSingleUint64(UnsignedLong.MAX_VALUE.longValue()) + .setSingleFloat(1.5f) + .setSingleDouble(2.5d) + .setSingleString("test") + .setSingleBytes(ByteString.copyFrom(new byte[] {0x01})) + .setSingleAny(Any.pack(DynamicMessage.newBuilder(BoolValue.of(true)).build())) + .setSingleDuration(com.google.protobuf.Duration.newBuilder().setSeconds(100)) + .setSingleTimestamp(Timestamp.newBuilder().setSeconds(100)) + .setSingleInt32Wrapper(Int32Value.of(5)) + .setSingleInt64Wrapper(Int64Value.of(10L)) + .setSingleUint32Wrapper(UInt32Value.of(1)) + .setSingleUint64Wrapper(UInt64Value.of(UnsignedLong.MAX_VALUE.longValue())) + .setSingleStringWrapper(StringValue.of("hello")) + .setSingleFloatWrapper(FloatValue.of(7.5f)) + .setSingleDoubleWrapper(DoubleValue.of(8.5d)) + .setSingleBytesWrapper(BytesValue.of(ByteString.copyFrom(new byte[] {0x02}))) + .addRepeatedInt64(5L) + .putMapStringString("a", "b") + .setStandaloneMessage(NestedMessage.getDefaultInstance()) + .setStandaloneEnum(NestedEnum.BAR) + .build(); + } }