From 43b0e670491f3229e4a7251a170fc413bda8444a Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Wed, 16 Sep 2026 17:31:26 -0700 Subject: [PATCH] Internal Changes PiperOrigin-RevId: 982819828 --- .../java/dev/cel/common/values/BUILD.bazel | 4 + .../values/ProtoLiteCelValueConverter.java | 453 ++++++++++++------ .../common/values/ProtoMessageLiteValue.java | 134 +++++- .../cel/common/values/ProtoMessageValue.java | 76 ++- .../values/RawProtoMessageLiteValue.java | 141 ++++-- .../cel/common/values/WireMessageLite.java | 47 ++ .../ProtoLiteCelValueConverterTest.java | 363 ++++++++++++-- .../values/ProtoMessageLiteValueTest.java | 190 ++++++++ .../common/values/ProtoMessageValueTest.java | 262 ++++++++-- .../values/RawProtoMessageLiteValueTest.java | 162 +++++-- .../CelLiteRuntimeVersionSkewTest.java | 300 +++++++++--- .../dev/cel/runtime/RuntimeEqualityTest.java | 37 +- 12 files changed, 1760 insertions(+), 409 deletions(-) create mode 100644 common/src/main/java/dev/cel/common/values/WireMessageLite.java 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 895a3410b..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", @@ -320,6 +322,7 @@ java_library( "ProtoLiteCelValueConverter.java", "ProtoMessageLiteValue.java", "RawProtoMessageLiteValue.java", + "WireMessageLite.java", ], tags = [ ], @@ -350,6 +353,7 @@ cel_android_library( "ProtoLiteCelValueConverter.java", "ProtoMessageLiteValue.java", "RawProtoMessageLiteValue.java", + "WireMessageLite.java", ], tags = [ ], diff --git a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java index 6aa523cb7..bdcb4cd6d 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java +++ b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java @@ -29,7 +29,6 @@ import com.google.protobuf.CodedInputStream; import com.google.protobuf.ExtensionRegistryLite; import com.google.protobuf.MessageLite; -import com.google.protobuf.MessageLiteOrBuilder; import com.google.protobuf.WireFormat; import dev.cel.common.annotations.Internal; import dev.cel.common.internal.CelLiteDescriptorPool; @@ -41,13 +40,12 @@ import java.io.IOException; import java.util.AbstractMap; import java.util.ArrayList; -import java.util.Collection; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; -import java.util.NoSuchElementException; import java.util.Optional; import java.util.TreeMap; +import org.jspecify.annotations.Nullable; /** * {@code ProtoLiteCelValueConverter} handles bidirectional conversion between native Java and @@ -62,8 +60,8 @@ @Immutable @Internal public final class ProtoLiteCelValueConverter extends BaseProtoCelValueConverter { - static final String MAP_KEY_FIELD_NAME = "key"; - static final String MAP_VALUE_FIELD_NAME = "value"; + private static final String MAP_KEY_FIELD_NAME = "key"; + private static final String MAP_VALUE_FIELD_NAME = "value"; private final CelLiteDescriptorPool descriptorPool; @@ -136,22 +134,18 @@ private static Object readFixed64BitField( } private Object readLengthDelimitedField( - CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException { + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { FieldLiteDescriptor.Type fieldType = fieldDescriptor.getProtoFieldType(); switch (fieldType) { case BYTES: return inputStream.readBytes(); case MESSAGE: - String fieldProtoTypeName = fieldDescriptor.getFieldProtoTypeName(); - MessageLiteDescriptor descriptor = - descriptorPool.findDescriptor(fieldProtoTypeName).orElse(null); - if (descriptor == null) { - return RawProtoMessageLiteValue.create(inputStream.readBytes(), fieldProtoTypeName, this); - } - MessageLite.Builder builder = descriptor.newMessageBuilder(); - inputStream.readMessage(builder, ExtensionRegistryLite.getEmptyRegistry()); - return builder.build(); + return mergeOrReadMessageField( + inputStream.readBytes(), fieldDescriptor.getFieldProtoTypeName(), existingValue); case STRING: return inputStream.readStringRequireUtf8(); default: @@ -159,6 +153,39 @@ private Object readLengthDelimitedField( } } + private Object readMessageField(ByteString bytes, String fieldProtoTypeName) { + return mergeOrReadMessageField(bytes, fieldProtoTypeName, /* existingValue= */ null); + } + + private Object mergeOrReadMessageField( + ByteString bytes, String fieldProtoTypeName, @Nullable Object existingValue) { + MessageLiteDescriptor descriptor = + descriptorPool.findDescriptor(fieldProtoTypeName).orElse(null); + if (descriptor == null) { + if (existingValue instanceof RawProtoMessageLiteValue) { + bytes = ((RawProtoMessageLiteValue) existingValue).toByteString().concat(bytes); + } + return RawProtoMessageLiteValue.create(bytes, fieldProtoTypeName, this); + } + WellKnownProto wellKnownProto = WellKnownProto.getByTypeName(fieldProtoTypeName).orElse(null); + if (isStructLike(wellKnownProto)) { + if (existingValue instanceof ProtoMessageLiteValue) { + bytes = ((ProtoMessageLiteValue) existingValue).toByteString().concat(bytes); + } + return ProtoMessageLiteValue.create(bytes, fieldProtoTypeName, this); + } + if (existingValue instanceof MessageLite) { + return mergeMessageLite(((MessageLite) existingValue).toBuilder(), bytes, fieldProtoTypeName); + } + return parseMessageLite(bytes, descriptor); + } + + // Unlike other WellKnownProtos (which unbox to CEL scalars/containers), FieldMask is + // represented as a standard struct message so field selection (e.g., mask.paths) works. + private static boolean isStructLike(@Nullable WellKnownProto wellKnownProto) { + return wellKnownProto == null || wellKnownProto == WellKnownProto.FIELD_MASK; + } + Object getDefaultCelValue(String protoTypeName, String fieldName) { MessageLiteDescriptor messageDescriptor = descriptorPool.getDescriptorOrThrow(protoTypeName); return getDefaultCelValue(messageDescriptor.getByFieldNameOrThrow(fieldName)); @@ -174,34 +201,43 @@ Optional findFieldDescriptor(String protoTypeName, int fiel .flatMap(desc -> desc.findByFieldNumber(fieldNumber)); } - Optional tryDecodeWellKnownProto(ByteString bytes, String protoTypeName) { - Optional wellKnownProto = WellKnownProto.getByTypeName(protoTypeName); - if (!wellKnownProto.isPresent()) { - return Optional.empty(); - } + MessageLite parseMessageLite(ByteString bytes, String protoTypeName) { + MessageLiteDescriptor descriptor = descriptorPool.getDescriptorOrThrow(protoTypeName); + return parseMessageLite(bytes, descriptor); + } - return descriptorPool - .findDescriptor(protoTypeName) - .map( - descriptor -> - decodeWellKnownProto(bytes, protoTypeName, descriptor, wellKnownProto.get())); + private static MessageLite parseMessageLite(ByteString bytes, MessageLiteDescriptor descriptor) { + if (bytes.isEmpty()) { + return descriptor.newMessageBuilder().getDefaultInstanceForType(); + } + return mergeMessageLite(descriptor.newMessageBuilder(), bytes, descriptor.getProtoTypeName()); } - private Object decodeWellKnownProto( - ByteString bytes, - String protoTypeName, - MessageLiteDescriptor descriptor, - WellKnownProto wellKnownProto) { + private static MessageLite mergeMessageLite( + MessageLite.Builder builder, ByteString bytes, String protoTypeName) { try { - MessageLite.Builder builder = descriptor.newMessageBuilder(); - builder.mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()); - return fromWellKnownProto(builder.build(), wellKnownProto); + return builder.mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()).build(); } catch (IOException e) { throw new IllegalArgumentException( - "Failed to decode well-known proto of type: " + protoTypeName, e); + "Failed to decode proto message of type: " + protoTypeName, e); } } + Optional tryDecodeProtoMessage(ByteString bytes, String protoTypeName) { + return descriptorPool + .findDescriptor(protoTypeName) + .map(descriptor -> decodeProtoMessage(bytes, protoTypeName, descriptor)); + } + + private Object decodeProtoMessage( + ByteString bytes, String protoTypeName, MessageLiteDescriptor descriptor) { + WellKnownProto wellKnownProto = WellKnownProto.getByTypeName(protoTypeName).orElse(null); + if (isStructLike(wellKnownProto)) { + return ProtoMessageLiteValue.create(bytes, protoTypeName, this); + } + return fromWellKnownProto(parseMessageLite(bytes, descriptor), checkNotNull(wellKnownProto)); + } + @Override public Object toRuntimeValue(Object value) { checkNotNull(value); @@ -212,35 +248,20 @@ public Object toRuntimeValue(Object value) { if (descriptor == null) { return RawProtoMessageLiteValue.create(msg.toByteString(), this); } - WellKnownProto wellKnownProto = - WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null); - - if (wellKnownProto == null) { - return ProtoMessageLiteValue.create(msg, descriptor.getProtoTypeName(), this); - } - - return fromWellKnownProto(msg, wellKnownProto); + return toRuntimeValue(msg, descriptor); } return super.toRuntimeValue(value); } - @Override - protected Object fromWellKnownProto(MessageLiteOrBuilder msg, WellKnownProto wellKnownProto) { - if (wellKnownProto == WellKnownProto.FIELD_MASK) { - MessageLite message = (MessageLite) msg; - MessageLiteDescriptor descriptor = - descriptorPool - .findDescriptor(message) - .orElseThrow( - () -> - new NoSuchElementException( - "Could not find a descriptor for message of type: " - + message.getClass().getName())); - return ProtoMessageLiteValue.create(message, descriptor.getProtoTypeName(), this); + private Object toRuntimeValue(MessageLite msg, MessageLiteDescriptor descriptor) { + WellKnownProto wellKnownProto = + WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null); + if (isStructLike(wellKnownProto)) { + return ProtoMessageLiteValue.create(msg, descriptor.getProtoTypeName(), this); } - return super.fromWellKnownProto(msg, wellKnownProto); + return fromWellKnownProto(msg, checkNotNull(wellKnownProto)); } private Object getDefaultValue(FieldLiteDescriptor fieldDescriptor) { @@ -286,141 +307,248 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) { if (WellKnownProto.isWrapperType(fieldProtoTypeName)) { return NullValue.NULL_VALUE; } - MessageLiteDescriptor descriptor = - descriptorPool.findDescriptor(fieldProtoTypeName).orElse(null); - if (descriptor == null) { - return RawProtoMessageLiteValue.create(ByteString.EMPTY, fieldProtoTypeName, this); - } - return descriptor.newMessageBuilder().build(); + return readMessageField(ByteString.EMPTY, fieldProtoTypeName); } throw new IllegalStateException("Unexpected java type: " + type); } - private ImmutableList readPackedRepeatedFields( + private Map.Entry readSingleMapEntry( CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException { + String entryTypeName = fieldDescriptor.getFieldProtoTypeName(); + MessageLiteDescriptor entryDescriptor = descriptorPool.getDescriptorOrThrow(entryTypeName); + FieldLiteDescriptor keyDescriptor = entryDescriptor.getByFieldNameOrThrow(MAP_KEY_FIELD_NAME); + FieldLiteDescriptor valueDescriptor = + entryDescriptor.getByFieldNameOrThrow(MAP_VALUE_FIELD_NAME); int length = inputStream.readInt32(); int oldLimit = inputStream.pushLimit(length); - ImmutableList.Builder builder = ImmutableList.builder(); + Object key = null; + Object value = null; while (inputStream.getBytesUntilLimit() > 0) { - builder.add(readPrimitiveField(inputStream, fieldDescriptor)); + int tag = inputStream.readTag(); + int tagWireType = WireFormat.getTagWireType(tag); + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber == keyDescriptor.getFieldNumber()) { + key = readSingularField(tagWireType, inputStream, keyDescriptor, key); + } else if (fieldNumber == valueDescriptor.getFieldNumber()) { + value = readSingularField(tagWireType, inputStream, valueDescriptor, value); + } else { + skipWireField(tag, inputStream); + } } inputStream.popLimit(oldLimit); - return builder.build(); - } - - private Map.Entry readSingleMapEntry( - CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException { - String entryTypeName = fieldDescriptor.getFieldProtoTypeName(); - ImmutableMap singleMapEntry = - readAllFields(inputStream.readByteArray(), entryTypeName).values(); - Object key = singleMapEntry.get(MAP_KEY_FIELD_NAME); if (key == null) { - key = getDefaultCelValue(entryTypeName, MAP_KEY_FIELD_NAME); + key = getDefaultCelValue(keyDescriptor); } - Object value = singleMapEntry.get(MAP_VALUE_FIELD_NAME); if (value == null) { - value = getDefaultCelValue(entryTypeName, MAP_VALUE_FIELD_NAME); + value = getDefaultCelValue(valueDescriptor); } - return new AbstractMap.SimpleEntry<>(key, value); + return new AbstractMap.SimpleImmutableEntry<>(key, value); + } + + @Nullable Object readSingleField(ByteString bytes, FieldLiteDescriptor fieldDescriptor) + throws IOException { + CodedInputStream inputStream = bytes.newCodedInput(); + int targetFieldNumber = fieldDescriptor.getFieldNumber(); + Object fieldValue = null; + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber != targetFieldNumber) { + skipWireField(tag, inputStream); + continue; + } + int tagWireType = WireFormat.getTagWireType(tag); + fieldValue = readFieldValue(tagWireType, inputStream, fieldDescriptor, fieldValue); + } + return fieldValue; + } + + boolean hasSingleField(ByteString bytes, FieldLiteDescriptor fieldDescriptor) throws IOException { + return hasSingleField( + bytes, + fieldDescriptor.getFieldNumber(), + fieldDescriptor.getEncodingType().equals(EncodingType.LIST) + && fieldDescriptor.getIsPacked()); + } + + static boolean hasSingleField(ByteString bytes, int targetFieldNumber, boolean isPackableList) + throws IOException { + CodedInputStream inputStream = bytes.newCodedInput(); + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber != targetFieldNumber) { + skipWireField(tag, inputStream); + continue; + } + int tagWireType = WireFormat.getTagWireType(tag); + // In protobuf wire format, a zero-length entry for a singular field (e.g. empty string, + // bytes, or empty submessage) represents explicit presence on the wire. Only packed + // repeated fields with empty payload represent an empty/absent collection. + if (isPackableList && tagWireType == WireFormat.WIRETYPE_LENGTH_DELIMITED) { + int length = inputStream.readInt32(); + inputStream.skipRawBytes(length); + if (length > 0) { + return true; + } + continue; + } + skipWireField(tag, inputStream); + return true; + } + return false; } - MessageFields readAllFields(byte[] bytes, String protoTypeName) throws IOException { + MessageFields readAllFields(ByteString bytes, String protoTypeName) throws IOException { MessageLiteDescriptor messageDescriptor = descriptorPool.getDescriptorOrThrow(protoTypeName); - CodedInputStream inputStream = CodedInputStream.newInstance(bytes); + if (bytes.isEmpty()) { + return MessageFields.EMPTY; + } + return readAllFields(bytes.newCodedInput(), messageDescriptor); + } - Multimap unknownFields = - Multimaps.newMultimap(new TreeMap<>(), ArrayList::new); - ImmutableMap.Builder fieldValues = ImmutableMap.builder(); - Map> repeatedFieldValues = new LinkedHashMap<>(); - Map> mapFieldValues = new LinkedHashMap<>(); + private MessageFields readAllFields( + CodedInputStream inputStream, MessageLiteDescriptor messageDescriptor) throws IOException { + Multimap unknownFields = null; + Map fieldValues = new LinkedHashMap<>(); for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { int tagWireType = WireFormat.getTagWireType(tag); int fieldNumber = WireFormat.getTagFieldNumber(tag); FieldLiteDescriptor fieldDescriptor = messageDescriptor.findByFieldNumber(fieldNumber).orElse(null); if (fieldDescriptor == null) { + if (unknownFields == null) { + unknownFields = Multimaps.newMultimap(new TreeMap<>(), ArrayList::new); + } unknownFields.put(fieldNumber, readUnknownField(tagWireType, inputStream)); continue; } - Object payload; - switch (tagWireType) { - case WireFormat.WIRETYPE_VARINT: - payload = readPrimitiveField(inputStream, fieldDescriptor); - break; - case WireFormat.WIRETYPE_FIXED32: - payload = readFixed32BitField(inputStream, fieldDescriptor); - break; - case WireFormat.WIRETYPE_FIXED64: - payload = readFixed64BitField(inputStream, fieldDescriptor); - break; - case WireFormat.WIRETYPE_LENGTH_DELIMITED: - EncodingType encodingType = fieldDescriptor.getEncodingType(); - switch (encodingType) { - case LIST: - if (fieldDescriptor.getIsPacked()) { - payload = readPackedRepeatedFields(inputStream, fieldDescriptor); - } else { - FieldLiteDescriptor.Type protoFieldType = fieldDescriptor.getProtoFieldType(); - boolean isLenDelimited = - protoFieldType.equals(FieldLiteDescriptor.Type.MESSAGE) - || protoFieldType.equals(FieldLiteDescriptor.Type.STRING) - || protoFieldType.equals(FieldLiteDescriptor.Type.BYTES); - if (!isLenDelimited) { - throw new IllegalStateException( - "Unexpected field type encountered for LEN-Delimited record: " - + protoFieldType); - } - - payload = readLengthDelimitedField(inputStream, fieldDescriptor); - } - break; - case MAP: - Map fieldMap = - mapFieldValues.computeIfAbsent(fieldNumber, (unused) -> new LinkedHashMap<>()); - Map.Entry mapEntry = readSingleMapEntry(inputStream, fieldDescriptor); - fieldMap.put(mapEntry.getKey(), mapEntry.getValue()); - payload = fieldMap; - break; - default: - payload = readLengthDelimitedField(inputStream, fieldDescriptor); - break; - } - break; - case WireFormat.WIRETYPE_START_GROUP: - case WireFormat.WIRETYPE_END_GROUP: - // TODO: Support groups - throw new UnsupportedOperationException("Groups are not supported"); - default: - throw new IllegalArgumentException("Unexpected wire type: " + tagWireType); + String fieldName = fieldDescriptor.getFieldName(); + Object fieldValue = + readFieldValue(tagWireType, inputStream, fieldDescriptor, fieldValues.get(fieldName)); + if (fieldValue != null) { + fieldValues.put(fieldName, fieldValue); } + } - if (fieldDescriptor.getEncodingType().equals(EncodingType.LIST)) { - String fieldName = fieldDescriptor.getFieldName(); - List repeatedValues = - repeatedFieldValues.computeIfAbsent(fieldNumber, (unused) -> new ArrayList<>()); + return MessageFields.create(ImmutableMap.copyOf(fieldValues), unknownFields); + } - if (payload instanceof Collection) { - repeatedValues.addAll((Collection) payload); - } else { - repeatedValues.add(payload); - } - if (!repeatedValues.isEmpty()) { - fieldValues.put(fieldName, repeatedValues); - } - } else { - fieldValues.put(fieldDescriptor.getFieldName(), payload); - } + private @Nullable Object readFieldValue( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + EncodingType encodingType = fieldDescriptor.getEncodingType(); + switch (encodingType) { + case SINGULAR: + return readSingularField(tagWireType, inputStream, fieldDescriptor, existingValue); + case LIST: + return readRepeatedField(tagWireType, inputStream, fieldDescriptor, existingValue); + case MAP: + return readMapField(tagWireType, inputStream, fieldDescriptor, existingValue); } + throw new IllegalStateException("Unexpected encoding type: " + encodingType); + } - // Protobuf encoding follows a "last one wins" semantics. This means for duplicated fields, - // we accept the last value encountered. - return MessageFields.create(fieldValues.buildKeepingLast(), unknownFields); + private Object readSingularField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + switch (tagWireType) { + case WireFormat.WIRETYPE_VARINT: + return readPrimitiveField(inputStream, fieldDescriptor); + case WireFormat.WIRETYPE_FIXED32: + return readFixed32BitField(inputStream, fieldDescriptor); + case WireFormat.WIRETYPE_FIXED64: + return readFixed64BitField(inputStream, fieldDescriptor); + case WireFormat.WIRETYPE_LENGTH_DELIMITED: + return readLengthDelimitedField(inputStream, fieldDescriptor, existingValue); + case WireFormat.WIRETYPE_START_GROUP: + case WireFormat.WIRETYPE_END_GROUP: + throw new UnsupportedOperationException("Groups are not supported"); + default: + throw new IllegalArgumentException("Unexpected wire type: " + tagWireType); + } + } + + // Safe because LIST fields only ever store an ArrayList as their accumulated value. + @SuppressWarnings("unchecked") + private @Nullable List readRepeatedField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + List repeatedValues = (List) existingValue; + if (tagWireType == WireFormat.WIRETYPE_LENGTH_DELIMITED && fieldDescriptor.getIsPacked()) { + return readPackedRepeatedFields(inputStream, fieldDescriptor, repeatedValues); + } + Object element = + readSingularField(tagWireType, inputStream, fieldDescriptor, /* existingValue= */ null); + if (repeatedValues == null) { + repeatedValues = new ArrayList<>(); + } + repeatedValues.add(element); + return repeatedValues; + } + + private static @Nullable List readPackedRepeatedFields( + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable List repeatedValues) + throws IOException { + int length = inputStream.readInt32(); + if (length == 0) { + return repeatedValues; + } + int oldLimit = inputStream.pushLimit(length); + if (repeatedValues == null) { + repeatedValues = new ArrayList<>(); + } + while (inputStream.getBytesUntilLimit() > 0) { + repeatedValues.add(readPrimitiveField(inputStream, fieldDescriptor)); + } + inputStream.popLimit(oldLimit); + return repeatedValues; } - MessageFields readMessageFields(MessageLite msg, String protoTypeName) throws IOException { - return readAllFields(msg.toByteArray(), protoTypeName); + // Safe because MAP fields only ever store a LinkedHashMap as their accumulated value. + @SuppressWarnings("unchecked") + private Map readMapField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + if (tagWireType != WireFormat.WIRETYPE_LENGTH_DELIMITED) { + throw new IllegalStateException("Unexpected wire type for map field: " + tagWireType); + } + Map mapValues = + existingValue != null ? (Map) existingValue : new LinkedHashMap<>(); + Map.Entry mapEntry = readSingleMapEntry(inputStream, fieldDescriptor); + mapValues.put(mapEntry.getKey(), mapEntry.getValue()); + return mapValues; + } + + static void skipWireField(int tag, CodedInputStream inputStream) throws IOException { + int tagWireType = WireFormat.getTagWireType(tag); + switch (tagWireType) { + case WireFormat.WIRETYPE_VARINT: + case WireFormat.WIRETYPE_FIXED64: + case WireFormat.WIRETYPE_LENGTH_DELIMITED: + case WireFormat.WIRETYPE_FIXED32: + inputStream.skipField(tag); + return; + case WireFormat.WIRETYPE_START_GROUP: + case WireFormat.WIRETYPE_END_GROUP: + throw new UnsupportedOperationException("Groups are not supported"); + default: + throw new IllegalArgumentException("Unknown wire type: " + tagWireType); + } } static Object readUnknownField(int tagWireType, CodedInputStream inputStream) throws IOException { @@ -447,19 +575,26 @@ static Object readUnknownField(int tagWireType, CodedInputStream inputStream) th @Immutable @SuppressWarnings("Immutable") // Safe immutable fields abstract static class MessageFields { + static final MessageFields EMPTY = + new AutoValue_ProtoLiteCelValueConverter_MessageFields( + ImmutableMap.of(), ImmutableListMultimap.of()); abstract ImmutableMap values(); abstract ImmutableListMultimap unknowns(); - static MessageFields create( - ImmutableMap fieldValues, Multimap unknownFields) { + private static MessageFields create( + ImmutableMap fieldValues, + @Nullable Multimap unknownFields) { return new AutoValue_ProtoLiteCelValueConverter_MessageFields( - fieldValues, ImmutableListMultimap.copyOf(unknownFields)); + fieldValues, + unknownFields == null + ? ImmutableListMultimap.of() + : ImmutableListMultimap.copyOf(unknownFields)); } } private ProtoLiteCelValueConverter(CelLiteDescriptorPool celLiteDescriptorPool) { - this.descriptorPool = celLiteDescriptorPool; + this.descriptorPool = checkNotNull(celLiteDescriptorPool); } } diff --git a/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java b/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java index 35acd5939..b1fe106db 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java @@ -21,12 +21,14 @@ import com.google.common.collect.ImmutableListMultimap; import com.google.common.collect.ImmutableMap; import com.google.errorprone.annotations.Immutable; +import com.google.protobuf.ByteString; import com.google.protobuf.MessageLite; import dev.cel.common.types.CelType; import dev.cel.common.types.StructTypeReference; import dev.cel.common.values.ProtoLiteCelValueConverter.MessageFields; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; import java.io.IOException; +import java.util.Objects; import java.util.Optional; import org.jspecify.annotations.Nullable; @@ -38,6 +40,11 @@ *

If the codebase has access to full protobuf messages with descriptors, use {@code * ProtoMessageValue} instead. * + *

An instance is backed by either a materialized {@link #rawValue()} or unparsed {@link + * #wireBytes()} (exactly one is non-null). Wire-backed instances decode selected fields directly + * from {@link #wireBytes()} and lazily parse the full {@link MessageLite} only if {@link #value()} + * is invoked. + * *

Implements {@link OptimizedSelectable} so that select chains can address fields by number: * *

    @@ -45,34 +52,60 @@ * resolving by {@link SelectField#fieldNumber()} maps the number to the runtime descriptor's * current field name, preventing {@code CelAttributeNotFoundException}. *
  • Version skew / unknown fields: When evaluating payloads serialized by a newer binary - * containing fields absent from the local {@code CelLiteDescriptor}, the unknown wire bytes - * are preserved in {@link #unknownFields()} and decoded on demand using the compile-time wire - * type and default metadata in {@link SelectField}. + * containing fields absent from the local {@code CelLiteDescriptor}, the unknown fields are + * decoded on demand directly from the message's wire bytes using the compile-time wire type + * and default metadata in {@link SelectField}. *
*/ @AutoValue @Immutable -public abstract class ProtoMessageLiteValue extends StructValue +abstract class ProtoMessageLiteValue extends StructValue implements OptimizedSelectable { - @Override - public abstract MessageLite value(); + // Populated when wrapping an already-materialized root MessageLite (e.g., from activation). + abstract @Nullable MessageLite rawValue(); + + // Populated when slicing a nested submessage from parent wire bytes to avoid deserializing and + // re-serializing intermediate hops; lazily parsed into a MessageLite only if value() is called. + abstract @Nullable ByteString wireBytes(); @Override public abstract CelType celType(); abstract ProtoLiteCelValueConverter protoLiteCelValueConverter(); + @Memoized + @Override + public MessageLite value() { + MessageLite msg = rawValue(); + if (msg != null) { + return msg; + } + return protoLiteCelValueConverter() + .parseMessageLite(checkNotNull(wireBytes()), celType().name()); + } + + @Memoized + ByteString serializedRawValue() { + return checkNotNull(rawValue()).toByteString(); + } + + ByteString toByteString() { + ByteString bytes = wireBytes(); + return bytes != null ? bytes : serializedRawValue(); + } + @Memoized MessageFields messageFields() { try { - return protoLiteCelValueConverter().readMessageFields(value(), celType().name()); + return protoLiteCelValueConverter().readAllFields(toByteString(), celType().name()); } catch (IOException e) { - throw new IllegalStateException("Unable to read message fields for " + celType().name(), e); + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); } } - ImmutableMap fieldValues() { + private ImmutableMap fieldValues() { return messageFields().values(); } @@ -82,9 +115,30 @@ ImmutableListMultimap unknownFields() { @Override public boolean isZeroValue() { + ByteString bytes = wireBytes(); + if (bytes != null && bytes.isEmpty()) { + return true; + } return value().getDefaultInstanceForType().equals(value()); } + @Override + public final boolean equals(Object other) { + if (other == this) { + return true; + } + if (!(other instanceof ProtoMessageLiteValue)) { + return false; + } + ProtoMessageLiteValue that = (ProtoMessageLiteValue) other; + return this.celType().equals(that.celType()) && this.value().equals(that.value()); + } + + @Override + public final int hashCode() { + return Objects.hash(value(), celType()); + } + @Override public Object select(String field) { return find(field) @@ -93,16 +147,15 @@ public Object select(String field) { @Override public Optional find(String field) { - Object fieldValue = fieldValues().get(field); - return Optional.ofNullable(fieldValue) - .map(value -> protoLiteCelValueConverter().toRuntimeValue(fieldValue)); + return Optional.ofNullable(fieldValues().get(field)) + .map(protoLiteCelValueConverter()::toRuntimeValue); } @Override public Object selectByFieldNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { - Object known = fieldValues().get(fd.getFieldName()); + Object known = readField(fd); if (known != null) { return protoLiteCelValueConverter().toRuntimeValue(known); } @@ -112,28 +165,48 @@ public Object selectByFieldNumber(SelectField field) { return protoLiteCelValueConverter().getDefaultCelValue(fd); } return RawProtoMessageLiteValue.selectWireOrDefault( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, + RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()), + protoLiteCelValueConverter()); } @Override public boolean hasFieldByNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { - return fieldValues().containsKey(fd.getFieldName()); + return hasField(fd); } - return RawProtoMessageLiteValue.isPresentInWire( - field, unknownFields().get(field.fieldNumber())); + return RawProtoMessageLiteValue.isPresentInWire(toByteString(), field); } @Override public Optional findByFieldNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { - return Optional.ofNullable(fieldValues().get(fd.getFieldName())) - .map(value -> protoLiteCelValueConverter().toRuntimeValue(value)); + return Optional.ofNullable(readField(fd)).map(protoLiteCelValueConverter()::toRuntimeValue); } return RawProtoMessageLiteValue.navigateWire( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, + RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()), + protoLiteCelValueConverter()); + } + + private @Nullable Object readField(FieldLiteDescriptor fd) { + try { + return protoLiteCelValueConverter().readSingleField(toByteString(), fd); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } + } + + private boolean hasField(FieldLiteDescriptor fd) { + try { + return protoLiteCelValueConverter().hasSingleField(toByteString(), fd); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } } private @Nullable FieldLiteDescriptor findFieldDescriptor(SelectField field) { @@ -142,13 +215,30 @@ public Optional findByFieldNumber(SelectField field) { .orElse(null); } - public static ProtoMessageLiteValue create( + static ProtoMessageLiteValue create( MessageLite value, String typeName, ProtoLiteCelValueConverter protoLiteCelValueConverter) { checkNotNull(value); checkNotNull(typeName); checkNotNull(protoLiteCelValueConverter); return new AutoValue_ProtoMessageLiteValue( - value, StructTypeReference.create(typeName), protoLiteCelValueConverter); + value, + /* wireBytes= */ null, + StructTypeReference.create(typeName), + protoLiteCelValueConverter); + } + + static ProtoMessageLiteValue create( + ByteString wireBytes, + String typeName, + ProtoLiteCelValueConverter protoLiteCelValueConverter) { + checkNotNull(wireBytes); + checkNotNull(typeName); + checkNotNull(protoLiteCelValueConverter); + return new AutoValue_ProtoMessageLiteValue( + /* rawValue= */ null, + wireBytes, + StructTypeReference.create(typeName), + protoLiteCelValueConverter); } ProtoMessageLiteValue() {} 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/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java index b0c41f058..e6bd6466b 100644 --- a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java @@ -30,7 +30,6 @@ import com.google.protobuf.ByteString; import com.google.protobuf.CodedInputStream; import com.google.protobuf.WireFormat; -import dev.cel.common.annotations.Internal; import dev.cel.common.exceptions.CelAttributeNotFoundException; import dev.cel.common.types.CelType; import dev.cel.common.types.StructTypeReference; @@ -39,6 +38,7 @@ import java.util.AbstractMap; import java.util.ArrayList; import java.util.List; +import java.util.Locale; import java.util.Map; import java.util.Optional; import java.util.TreeMap; @@ -49,7 +49,7 @@ * client-server version skew issues where newer fields or submessages lack generated classes and * descriptors in the evaluation environment. * - *

Rather than requiring compiled {@link MessageLite} classes or runtime schema descriptors, this + *

Rather than requiring compiled {@code MessageLite} classes or runtime schema descriptors, this * value encapsulates the raw wire-format {@link ByteString} payload and performs classless, * reflection-free field traversal directly over wire tags via {@link CodedInputStream}. */ @@ -57,15 +57,15 @@ @AutoValue.CopyAnnotations @Immutable @SuppressWarnings("Immutable") // Immutable wire fields -@Internal -public abstract class RawProtoMessageLiteValue extends StructValue - implements OptimizedSelectable { +abstract class RawProtoMessageLiteValue extends StructValue + implements OptimizedSelectable, WireMessageLite { private static final String UNKNOWN_MESSAGE_TYPE_NAME = "cel.@unknownMessage"; private static final int MAP_KEY_FIELD_NUMBER = 1; private static final int MAP_VALUE_FIELD_NUMBER = 2; - abstract ByteString rawWireBytes(); + @Override + public abstract ByteString toByteString(); @Override public abstract CelType celType(); @@ -73,14 +73,39 @@ public abstract class RawProtoMessageLiteValue extends StructValue unknownFields() { try { - CodedInputStream inputStream = rawWireBytes().newCodedInput(); + CodedInputStream inputStream = toByteString().newCodedInput(); Multimap fields = Multimaps.newMultimap(new TreeMap<>(), ArrayList::new); for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { int tagWireType = WireFormat.getTagWireType(tag); @@ -96,7 +121,7 @@ ImmutableListMultimap unknownFields() { @Override public boolean isZeroValue() { - return rawWireBytes().isEmpty(); + return toByteString().isEmpty(); } /** @@ -128,18 +153,45 @@ private CelAttributeNotFoundException newUnoptimizedFieldResolutionException(Str @Override public Object selectByFieldNumber(SelectField field) { return selectWireOrDefault( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, readWireField(toByteString(), field.fieldNumber()), protoLiteCelValueConverter()); } @Override public boolean hasFieldByNumber(SelectField field) { - return isPresentInWire(field, unknownFields().get(field.fieldNumber())); + return isPresentInWire(toByteString(), field); } @Override public Optional findByFieldNumber(SelectField field) { return navigateWire( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, readWireField(toByteString(), field.fieldNumber()), protoLiteCelValueConverter()); + } + + /** + * Scans {@code wireBytes} for a single {@code targetFieldNumber}, skipping all other wire tags. + * + *

Package-private: shared with {@code ProtoMessageLiteValue} for unknown field resolution. + */ + static ImmutableList readWireField(ByteString wireBytes, int targetFieldNumber) { + if (wireBytes.isEmpty()) { + return ImmutableList.of(); + } + ImmutableList.Builder entries = ImmutableList.builder(); + try { + CodedInputStream inputStream = wireBytes.newCodedInput(); + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber != targetFieldNumber) { + ProtoLiteCelValueConverter.skipWireField(tag, inputStream); + continue; + } + int tagWireType = WireFormat.getTagWireType(tag); + entries.add(ProtoLiteCelValueConverter.readUnknownField(tagWireType, inputStream)); + } + } catch (IOException e) { + throw new IllegalStateException("Failed to parse raw proto message wire bytes", e); + } + return entries.build(); } /** @@ -169,55 +221,40 @@ private static Object decodeWireField( } boolean isRepeated = field.defaultValue() instanceof List; - String protoTypeName = resolveProtoTypeName(field); - return decodeWireEntries(unknowns, typeCode, protoTypeName, isRepeated, converter); - } - - /** - * Resolves the protobuf message type name for a field from the optimizer metadata in {@link - * SelectField}, or {@link #UNKNOWN_MESSAGE_TYPE_NAME} if unspecified. - */ - private static String resolveProtoTypeName(SelectField field) { - if (!field.protoTypeName().isEmpty()) { - return field.protoTypeName(); - } - return UNKNOWN_MESSAGE_TYPE_NAME; + return decodeWireEntries(unknowns, typeCode, field.protoTypeName(), isRepeated, converter); } private static Object resolveDefault(SelectField field, ProtoLiteCelValueConverter converter) { if (field.defaultValue() != null) { return field.defaultValue(); } - return create(ByteString.EMPTY, resolveProtoTypeName(field), converter); + return decodeMessageValue(ByteString.EMPTY, field.protoTypeName(), converter); } /** - * Returns whether a field has presence in preserved wire bytes. + * Scans {@code wireBytes} to determine whether {@code field} is present on the wire. * *

Package-private: shared with {@code ProtoMessageLiteValue} for unknown field resolution. */ - static boolean isPresentInWire(SelectField field, ImmutableList unknowns) { + static boolean isPresentInWire(ByteString wireBytes, SelectField field) { + try { + return ProtoLiteCelValueConverter.hasSingleField( + wireBytes, field.fieldNumber(), isPackableRepeated(field)); + } catch (IOException e) { + throw new IllegalStateException("Failed to parse raw proto message wire bytes", e); + } + } + + private static boolean isPresentInWire(SelectField field, ImmutableList unknowns) { if (unknowns.isEmpty()) { return false; } - boolean isRepeated = field.defaultValue() instanceof List; - int typeCode = field.typeCode(); - // In protobuf wire format, a zero-length entry for a singular field (e.g. empty string, // bytes, or empty submessage) represents explicit presence on the wire. Only packed repeated // fields with empty payload represent an empty/absent collection. - if (!isRepeated) { - return true; - } - - boolean isPackable = - typeCode != FieldLiteDescriptor.Type.STRING.getNumber() - && typeCode != FieldLiteDescriptor.Type.BYTES.getNumber() - && typeCode != FieldLiteDescriptor.Type.MESSAGE.getNumber() - && typeCode != FieldLiteDescriptor.Type.GROUP.getNumber(); - if (!isPackable) { + if (!isPackableRepeated(field)) { return true; } @@ -229,6 +266,13 @@ static boolean isPresentInWire(SelectField field, ImmutableList unknowns return false; } + private static boolean isPackableRepeated(SelectField field) { + return (field.defaultValue() instanceof List) + && FieldLiteDescriptor.Type.forNumber(field.typeCode()) + .toWireFormatFieldType() + .isPackable(); + } + /** * Navigates a field on preserved wire bytes, returning empty if absent. * @@ -419,9 +463,7 @@ static Object decodeWireValue( throw new UnsupportedOperationException("Groups are not supported"); case MESSAGE: ByteString msgBytes = requireType(raw, ByteString.class, fieldType); - return converter - .tryDecodeWellKnownProto(msgBytes, protoTypeName) - .orElseGet(() -> create(msgBytes, protoTypeName, converter)); + return decodeMessageValue(msgBytes, protoTypeName, converter); case BYTES: return CelByteString.of(requireType(raw, ByteString.class, fieldType).toByteArray()); case UINT32: @@ -437,6 +479,13 @@ static Object decodeWireValue( throw new IllegalArgumentException("Unsupported proto field type: " + fieldType); } + private static Object decodeMessageValue( + ByteString msgBytes, String protoTypeName, ProtoLiteCelValueConverter converter) { + return converter + .tryDecodeProtoMessage(msgBytes, protoTypeName) + .orElseGet(() -> create(msgBytes, protoTypeName, converter)); + } + private static T requireType( Object raw, Class expectedType, WireFormat.FieldType fieldType) { if (!expectedType.isInstance(raw)) { @@ -509,12 +558,12 @@ private static ImmutableList decodePacked( } } - public static RawProtoMessageLiteValue create( + static RawProtoMessageLiteValue create( ByteString rawWireBytes, ProtoLiteCelValueConverter protoLiteCelValueConverter) { return create(rawWireBytes, "", protoLiteCelValueConverter); } - public static RawProtoMessageLiteValue create( + static RawProtoMessageLiteValue create( ByteString rawWireBytes, String protoTypeName, ProtoLiteCelValueConverter protoLiteCelValueConverter) { diff --git a/common/src/main/java/dev/cel/common/values/WireMessageLite.java b/common/src/main/java/dev/cel/common/values/WireMessageLite.java new file mode 100644 index 000000000..eb0572670 --- /dev/null +++ b/common/src/main/java/dev/cel/common/values/WireMessageLite.java @@ -0,0 +1,47 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dev.cel.common.values; + +import com.google.errorprone.annotations.Immutable; +import com.google.protobuf.ByteString; +import dev.cel.common.annotations.Beta; + +/** + * Represents a protobuf message evaluation result in {@code CelLiteRuntime} when no {@code + * CelLiteDescriptor} is registered for the message type. + * + *

When a message-typed expression is evaluated in {@code CelLiteRuntime}: + * + *

    + *
  • If a {@code CelLiteDescriptor} is registered for the message type, evaluation produces a + * {@code MessageLite} instance. + *
  • Otherwise, evaluation produces a {@code WireMessageLite} carrying the message's protobuf + * type name and wire-encoded payload. + *
+ */ +@Immutable +@Beta +public interface WireMessageLite { + + /** + * Returns the fully-qualified protobuf message type name (e.g. {@code + * "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"}), or {@code "cel.@unknownMessage"} if + * the message type name is not known at runtime. + */ + String protoTypeName(); + + /** Serializes the message to a {@link ByteString} in protobuf wire format. */ + ByteString toByteString(); +} diff --git a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java index 077068ce8..940842567 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java @@ -19,12 +19,14 @@ import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableListMultimap; +import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.google.common.collect.Multimap; import com.google.common.primitives.UnsignedLong; import com.google.protobuf.BoolValue; import com.google.protobuf.ByteString; import com.google.protobuf.BytesValue; +import com.google.protobuf.CodedOutputStream; import com.google.protobuf.DoubleValue; import com.google.protobuf.Duration; import com.google.protobuf.ExtensionRegistryLite; @@ -38,6 +40,7 @@ import com.google.protobuf.Timestamp; import com.google.protobuf.UInt32Value; import com.google.protobuf.UInt64Value; +import com.google.protobuf.WireFormat; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.common.internal.CelLiteDescriptorPool; @@ -45,9 +48,11 @@ import dev.cel.common.values.ProtoLiteCelValueConverter.MessageFields; import dev.cel.expr.conformance.proto3.NestedTestAllTypes; import dev.cel.expr.conformance.proto3.TestAllTypes; +import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage; import dev.cel.expr.conformance.proto3.TestAllTypesCelDescriptor; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; import dev.cel.protobuf.CelLiteDescriptor.MessageLiteDescriptor; +import java.io.ByteArrayOutputStream; import java.io.IOException; import java.time.Instant; import java.util.LinkedHashMap; @@ -57,7 +62,7 @@ import org.junit.runner.RunWith; @RunWith(TestParameterInjector.class) -public class ProtoLiteCelValueConverterTest { +public final class ProtoLiteCelValueConverterTest { private static final CelLiteDescriptorPool EMPTY_DESCRIPTOR_POOL = new CelLiteDescriptorPool() { @Override @@ -101,9 +106,10 @@ public void fromProtoMessageToCelValue_withoutDescriptor_returnsRawProtoMessageL Object adaptedValue = converterWithoutDescriptors.toRuntimeValue(msg); - assertThat(adaptedValue) - .isEqualTo( - RawProtoMessageLiteValue.create(msg.toByteString(), converterWithoutDescriptors)); + assertThat(adaptedValue).isInstanceOf(RawProtoMessageLiteValue.class); + RawProtoMessageLiteValue rawValue = (RawProtoMessageLiteValue) adaptedValue; + assertThat(rawValue.toByteString()).isEqualTo(msg.toByteString()); + assertThat(rawValue.protoTypeName()).isEqualTo("cel.@unknownMessage"); } @SuppressWarnings("ImmutableEnumChecker") // Test only @@ -174,7 +180,7 @@ public void readAllFields_repeatedFields_packedBytesCombinations( @TestParameter RepeatedFieldBytesTestCase testCase) throws Exception { MessageFields fields = PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( - testCase.bytes, "cel.expr.conformance.proto3.TestAllTypes"); + ByteString.copyFrom(testCase.bytes), "cel.expr.conformance.proto3.TestAllTypes"); assertThat(fields.values()).containsExactly("repeated_int64", ImmutableList.of(1L, 2L, 3L)); } @@ -258,7 +264,7 @@ public void unknowns_repeatedEncodedBytes_allRecordsKeptWithKeysSorted() throws MessageFields messageFields = PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( - bytes, "cel.expr.conformance.proto3.TestAllTypes"); + ByteString.copyFrom(bytes), "cel.expr.conformance.proto3.TestAllTypes"); assertThat(messageFields.values()).isEmpty(); assertThat(messageFields.unknowns()) @@ -275,7 +281,7 @@ public void readAllFields_unknownFields(@TestParameter UnknownFieldsTestCase tes MessageFields messageFields = PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( - testCase.bytes, "cel.expr.conformance.proto3.TestAllTypes"); + ByteString.copyFrom(testCase.bytes), "cel.expr.conformance.proto3.TestAllTypes"); assertThat(messageFields.values()).isEmpty(); assertThat(messageFields.unknowns()).containsExactlyEntriesIn(testCase.unknownMap).inOrder(); @@ -316,7 +322,7 @@ public void readAllFields_unknownFieldsWithValues() throws Exception { MessageFields fields = PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( - unknownMessageBytes, "cel.expr.conformance.proto3.TestAllTypes"); + ByteString.copyFrom(unknownMessageBytes), "cel.expr.conformance.proto3.TestAllTypes"); assertThat(TextFormat.printer().printToString(parsedMsg)) .isEqualTo( @@ -388,12 +394,11 @@ public void getDefaultCelValue_nestedMessageWithoutDescriptor_returnsRawProtoMes Object defaultValue = converterWithoutNested.getDefaultCelValue(nestedMsgField); - assertThat(defaultValue) - .isEqualTo( - RawProtoMessageLiteValue.create( - ByteString.EMPTY, - "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", - converterWithoutNested)); + assertThat(defaultValue).isInstanceOf(RawProtoMessageLiteValue.class); + RawProtoMessageLiteValue rawValue = (RawProtoMessageLiteValue) defaultValue; + assertThat(rawValue.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(rawValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); } @Test @@ -409,74 +414,360 @@ public void readAllFields_nestedMessageWithoutDescriptor_returnsRawProtoMessageL MessageFields fields = PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( - msg.toByteArray(), "cel.expr.conformance.proto3.TestAllTypes"); - - assertThat(fields.values()) - .containsExactly( - "oneof_type", - RawProtoMessageLiteValue.create( - msg.getOneofType().toByteString(), - "cel.expr.conformance.proto3.NestedTestAllTypes", - PROTO_LITE_CEL_VALUE_CONVERTER)); + msg.toByteString(), "cel.expr.conformance.proto3.TestAllTypes"); + + assertThat(fields.values().keySet()).containsExactly("oneof_type"); + Object fieldValue = fields.values().get("oneof_type"); + assertThat(fieldValue).isInstanceOf(RawProtoMessageLiteValue.class); + RawProtoMessageLiteValue rawValue = (RawProtoMessageLiteValue) fieldValue; + assertThat(rawValue.toByteString()).isEqualTo(msg.getOneofType().toByteString()); + assertThat(rawValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } @Test - public void tryDecodeWellKnownProto_validBytes_returnsDecodedValue() { + public void tryDecodeProtoMessage_wellKnownType_returnsDecodedValue() { Int32Value int32Value = Int32Value.of(42); Optional decoded = - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( int32Value.toByteString(), "google.protobuf.Int32Value"); assertThat(decoded).hasValue(42L); } @Test - public void tryDecodeWellKnownProto_notWellKnownType_returnsEmpty() { + public void tryDecodeProtoMessage_registeredMessageType_returnsWireBackedProtoMessageLiteValue() { + NestedMessage nestedMsg = NestedMessage.newBuilder().setBb(42).build(); + Optional decoded = - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( - ByteString.EMPTY, "cel.expr.conformance.proto3.TestAllTypes"); + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( + nestedMsg.toByteString(), "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + assertThat(decoded.map(v -> ((ProtoMessageLiteValue) v).rawValue())).isEmpty(); + assertThat(decoded.map(v -> ((ProtoMessageLiteValue) v).wireBytes())) + .hasValue(nestedMsg.toByteString()); + assertThat(decoded) + .hasValue( + ProtoMessageLiteValue.create( + nestedMsg, + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", + PROTO_LITE_CEL_VALUE_CONVERTER)); + } - assertThat(decoded).isEmpty(); + @Test + public void tryDecodeProtoMessage_fieldMask_returnsWireBackedProtoMessageLiteValue() { + FieldMask fieldMask = FieldMask.newBuilder().addPaths("foo").addPaths("bar").build(); + + Optional decoded = + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( + fieldMask.toByteString(), "google.protobuf.FieldMask"); + + assertThat(decoded.map(v -> ((ProtoMessageLiteValue) v).rawValue())).isEmpty(); + assertThat(decoded.map(v -> ((ProtoMessageLiteValue) v).select("paths"))) + .hasValue(ImmutableList.of("foo", "bar")); } @Test - public void tryDecodeWellKnownProto_missingDescriptor_returnsEmpty() { + public void + tryDecodeProtoMessage_registeredMessageTypeEmptyBytes_returnsDefaultProtoMessageLiteValue() { + Optional decoded = + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( + ByteString.EMPTY, "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + assertThat(decoded) + .hasValue( + ProtoMessageLiteValue.create( + NestedMessage.getDefaultInstance(), + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", + PROTO_LITE_CEL_VALUE_CONVERTER)); + } + + @Test + public void tryDecodeProtoMessage_missingDescriptor_returnsEmpty() { ProtoLiteCelValueConverter converter = ProtoLiteCelValueConverter.newInstance(EMPTY_DESCRIPTOR_POOL); Optional decoded = - converter.tryDecodeWellKnownProto(ByteString.EMPTY, "google.protobuf.Int32Value"); + converter.tryDecodeProtoMessage(ByteString.EMPTY, "google.protobuf.Int32Value"); assertThat(decoded).isEmpty(); } @Test - public void tryDecodeWellKnownProto_invalidBytes_throwsIllegalArgumentException() { + public void tryDecodeProtoMessage_invalidBytes_throwsIllegalArgumentException() { ByteString corruptBytes = ByteString.copyFrom(new byte[] {(byte) 0xFF, (byte) 0xFF}); IllegalArgumentException exception = assertThrows( IllegalArgumentException.class, () -> - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( corruptBytes, "google.protobuf.Int32Value")); assertThat(exception) .hasMessageThat() - .contains("Failed to decode well-known proto of type: google.protobuf.Int32Value"); + .contains("Failed to decode proto message of type: google.protobuf.Int32Value"); assertThat(exception).hasCauseThat().isInstanceOf(IOException.class); } @Test - public void tryDecodeWellKnownProto_anyType_throwsUnsupportedOperationException() { + public void tryDecodeProtoMessage_anyType_throwsUnsupportedOperationException() { UnsupportedOperationException exception = assertThrows( UnsupportedOperationException.class, () -> - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( ByteString.EMPTY, "google.protobuf.Any")); assertThat(exception).hasMessageThat().contains("ANY_VALUE"); } + + @Test + public void readAllFields_splitSingularSubmessages_mergesAllOccurrences() throws Exception { + ByteArrayOutputStream unknownFieldBaos = new ByteArrayOutputStream(); + CodedOutputStream cos = CodedOutputStream.newInstance(unknownFieldBaos); + cos.writeInt64(999, 42L); + cos.flush(); + NestedMessage nestedWithUnknown = + NestedMessage.parseFrom( + unknownFieldBaos.toByteArray(), ExtensionRegistryLite.getEmptyRegistry()); + TestAllTypes part1 = + TestAllTypes.newBuilder() + .setOneofType( + NestedTestAllTypes.newBuilder() + .setPayload(TestAllTypes.newBuilder().setSingleInt32(10))) + .setSingleDuration(Duration.newBuilder().setSeconds(10)) + .setSingleNestedMessage(NestedMessage.newBuilder().setBb(99)) + .build(); + TestAllTypes part2 = + TestAllTypes.newBuilder() + .setOneofType( + NestedTestAllTypes.newBuilder() + .setPayload(TestAllTypes.newBuilder().setSingleString("merged"))) + .setSingleDuration(Duration.newBuilder().setNanos(500)) + .setSingleNestedMessage(nestedWithUnknown) + .build(); + ByteString splitWireBytes = part1.toByteString().concat(part2.toByteString()); + + MessageFields fields = + PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( + splitWireBytes, "cel.expr.conformance.proto3.TestAllTypes"); + + assertThat(fields.values().get("single_duration")) + .isEqualTo(Duration.newBuilder().setSeconds(10).setNanos(500).build()); + ProtoMessageLiteValue nestedMsg = + (ProtoMessageLiteValue) fields.values().get("single_nested_message"); + assertThat(nestedMsg.rawValue()).isNull(); + assertThat(nestedMsg.select("bb")).isEqualTo(99L); + assertThat(nestedMsg.unknownFields()).valuesForKey(999).containsExactly(42L); + RawProtoMessageLiteValue rawSubmessage = + (RawProtoMessageLiteValue) fields.values().get("oneof_type"); + assertThat( + NestedTestAllTypes.parseFrom( + rawSubmessage.toByteString(), ExtensionRegistryLite.getEmptyRegistry())) + .isEqualTo( + NestedTestAllTypes.newBuilder() + .setPayload(TestAllTypes.newBuilder().setSingleInt32(10).setSingleString("merged")) + .build()); + } + + @Test + public void parseMessageLite_emptyBytes_returnsDefaultInstanceSingleton() { + MessageLite parsed = + PROTO_LITE_CEL_VALUE_CONVERTER.parseMessageLite( + ByteString.EMPTY, "cel.expr.conformance.proto3.TestAllTypes"); + + assertThat(parsed).isSameInstanceAs(TestAllTypes.getDefaultInstance()); + } + + @Test + public void readAllFields_emptyBytes_returnsEmptySingleton() throws Exception { + MessageFields fields = + PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( + ByteString.EMPTY, "cel.expr.conformance.proto3.TestAllTypes"); + + assertThat(fields).isSameInstanceAs(MessageFields.EMPTY); + } + + @Test + public void readSingleField_emptyBytes_returnsNull() throws Exception { + FieldLiteDescriptor fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", TestAllTypes.SINGLE_INT64_FIELD_NUMBER) + .get(); + + Object result = PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(ByteString.EMPTY, fd); + + assertThat(result).isNull(); + } + + @Test + public void readSingleField_absentField_returnsNull() throws Exception { + TestAllTypes proto = TestAllTypes.newBuilder().setSingleInt64(42L).build(); + FieldLiteDescriptor fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", TestAllTypes.SINGLE_BOOL_FIELD_NUMBER) + .get(); + + Object result = PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(proto.toByteString(), fd); + + assertThat(result).isNull(); + } + + @Test + public void readSingleField_emptyPackedRepeated_returnsNull() throws Exception { + ByteArrayOutputStream emptyPackedOut = new ByteArrayOutputStream(); + CodedOutputStream emptyPackedCos = CodedOutputStream.newInstance(emptyPackedOut); + emptyPackedCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + emptyPackedCos.flush(); + ByteString emptyPackedBytes = ByteString.copyFrom(emptyPackedOut.toByteArray()); + FieldLiteDescriptor repeatedInt32Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", + TestAllTypes.REPEATED_INT32_FIELD_NUMBER) + .get(); + + Object result = + PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(emptyPackedBytes, repeatedInt32Fd); + + assertThat(result).isNull(); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum ReadSingleFieldTestCase { + SINGLE_STRING(TestAllTypes.SINGLE_STRING_FIELD_NUMBER, "target_str"), + REPEATED_STRING(TestAllTypes.REPEATED_STRING_FIELD_NUMBER, ImmutableList.of("a", "b")), + REPEATED_INT32(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, ImmutableList.of(1, 2)), + MAP_STRING_STRING( + TestAllTypes.MAP_STRING_STRING_FIELD_NUMBER, + ImmutableMap.of("k", "v", "k2", "v2", "", "")), + SPLIT_DURATION( + TestAllTypes.SINGLE_DURATION_FIELD_NUMBER, + Duration.newBuilder().setSeconds(10).setNanos(500).build()); + + private final int fieldNumber; + private final Object expected; + + ReadSingleFieldTestCase(int fieldNumber, Object expected) { + this.fieldNumber = fieldNumber; + this.expected = expected; + } + } + + @Test + public void readSingleField_skipsOtherFieldsAndDecodesTarget( + @TestParameter ReadSingleFieldTestCase testCase) throws Exception { + TestAllTypes part1 = + TestAllTypes.newBuilder() + .setSingleInt64(42L) + .setSingleFixed32(10) + .setSingleFixed64(20L) + .setSingleString("target_str") + .addRepeatedString("a") + .addRepeatedInt32(1) + .putMapStringString("k", "v") + .setSingleDuration(Duration.newBuilder().setSeconds(10)) + .build(); + TestAllTypes part2 = + TestAllTypes.newBuilder() + .addRepeatedString("b") + .addRepeatedInt32(2) + .putMapStringString("k2", "v2") + .setSingleDuration(Duration.newBuilder().setNanos(500)) + .build(); + ByteArrayOutputStream mapEntryWithUnknownOut = new ByteArrayOutputStream(); + CodedOutputStream mapEntryWithUnknownCos = CodedOutputStream.newInstance(mapEntryWithUnknownOut); + mapEntryWithUnknownCos.writeInt64(3, 99L); + mapEntryWithUnknownCos.flush(); + ByteArrayOutputStream extraWireOut = new ByteArrayOutputStream(); + CodedOutputStream extraWireCos = CodedOutputStream.newInstance(extraWireOut); + extraWireCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + extraWireCos.writeByteArray( + TestAllTypes.MAP_STRING_STRING_FIELD_NUMBER, mapEntryWithUnknownOut.toByteArray()); + extraWireCos.flush(); + ByteString bytes = + part1 + .toByteString() + .concat(part2.toByteString()) + .concat(ByteString.copyFrom(extraWireOut.toByteArray())); + FieldLiteDescriptor fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor("cel.expr.conformance.proto3.TestAllTypes", testCase.fieldNumber) + .get(); + + Object result = PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(bytes, fd); + + assertThat(result).isEqualTo(testCase.expected); + } + + @Test + public void hasSingleField_emptyBytes_returnsFalse() throws Exception { + FieldLiteDescriptor singleInt64Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", TestAllTypes.SINGLE_INT64_FIELD_NUMBER) + .get(); + + boolean result = PROTO_LITE_CEL_VALUE_CONVERTER.hasSingleField(ByteString.EMPTY, singleInt64Fd); + + assertThat(result).isFalse(); + } + + @Test + public void hasSingleField_emptyPackedRepeatedField_returnsFalse() throws Exception { + FieldLiteDescriptor repeatedInt32Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", + TestAllTypes.REPEATED_INT32_FIELD_NUMBER) + .get(); + ByteArrayOutputStream emptyPackedOut = new ByteArrayOutputStream(); + CodedOutputStream emptyPackedCos = CodedOutputStream.newInstance(emptyPackedOut); + emptyPackedCos.writeInt64(TestAllTypes.SINGLE_INT64_FIELD_NUMBER, 42L); + emptyPackedCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + emptyPackedCos.flush(); + ByteString emptyPackedBytes = ByteString.copyFrom(emptyPackedOut.toByteArray()); + + boolean result = + PROTO_LITE_CEL_VALUE_CONVERTER.hasSingleField(emptyPackedBytes, repeatedInt32Fd); + + assertThat(result).isFalse(); + } + + @Test + public void hasSingleField_emptyPackedFollowedByPopulatedPacked_returnsTrue() throws Exception { + FieldLiteDescriptor repeatedInt32Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", + TestAllTypes.REPEATED_INT32_FIELD_NUMBER) + .get(); + ByteArrayOutputStream emptyThenPopulatedOut = new ByteArrayOutputStream(); + CodedOutputStream emptyThenPopulatedCos = CodedOutputStream.newInstance(emptyThenPopulatedOut); + emptyThenPopulatedCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + emptyThenPopulatedCos.writeByteArray( + TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[] {1, 2}); + emptyThenPopulatedCos.flush(); + ByteString emptyThenPopulatedBytes = ByteString.copyFrom(emptyThenPopulatedOut.toByteArray()); + + boolean result = + PROTO_LITE_CEL_VALUE_CONVERTER.hasSingleField(emptyThenPopulatedBytes, repeatedInt32Fd); + + assertThat(result).isTrue(); + } + + @Test + public void skipWireField_groupWireType_throwsUnsupportedOperationException() { + int startGroupTag = (1 << 3) | WireFormat.WIRETYPE_START_GROUP; + + assertThrows( + UnsupportedOperationException.class, + () -> + ProtoLiteCelValueConverter.skipWireField( + startGroupTag, ByteString.EMPTY.newCodedInput())); + } } diff --git a/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java b/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java index e7d7ce4de..41c9fe666 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java @@ -21,6 +21,7 @@ import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.google.common.primitives.UnsignedLong; +import com.google.common.testing.EqualsTester; import com.google.protobuf.Any; import com.google.protobuf.BoolValue; import com.google.protobuf.ByteString; @@ -46,6 +47,7 @@ import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage; import dev.cel.expr.conformance.proto3.TestAllTypesCelDescriptor; import java.io.ByteArrayOutputStream; +import java.io.IOException; import java.time.Duration; import java.time.Instant; import java.util.Optional; @@ -86,6 +88,186 @@ public void create_withPopulatedMessage() { assertThat(messageLiteValue.isZeroValue()).isFalse(); } + @Test + public void create_withEmptyByteString() { + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + ByteString.EMPTY, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + + assertThat(messageLiteValue.isZeroValue()).isTrue(); + assertThat(messageLiteValue.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(messageLiteValue.value()).isSameInstanceAs(TestAllTypes.getDefaultInstance()); + } + + @Test + public void isZeroValue_emptyWireBytesWithoutDescriptor_returnsTrueWithoutParsing() { + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + ByteString.EMPTY, "unregistered.Message", PROTO_LITE_CEL_VALUE_CONVERTER); + + assertThat(messageLiteValue.isZeroValue()).isTrue(); + } + + @Test + public void isZeroValue_nonEmptyWireBytesForDefaultMessage_returnsTrue() { + // Explicit wire tag for field 1 (single_int32) with value 0: deserializes to default instance. + ByteString explicitZeroFieldBytes = ByteString.copyFrom(new byte[] {0x08, 0x00}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + explicitZeroFieldBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + + assertThat(messageLiteValue.isZeroValue()).isTrue(); + } + + @Test + public void create_withPopulatedByteString_selectsAndLazilyMaterializesValue() { + TestAllTypes expected = + TestAllTypes.newBuilder().setSingleInt64(42L).setSingleString("hello").build(); + + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + expected.toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + + assertThat(messageLiteValue.select("single_int64")).isEqualTo(42L); + assertThat(messageLiteValue.select("single_string")).isEqualTo("hello"); + assertThat(messageLiteValue.isZeroValue()).isFalse(); + assertThat(messageLiteValue.toByteString()).isEqualTo(expected.toByteString()); + assertThat(messageLiteValue.value()).isEqualTo(expected); + } + + @Test + public void equals_byteStringBackedAndMessageBacked_areEqual() { + TestAllTypes populated = TestAllTypes.newBuilder().setSingleInt64(42L).build(); + TestAllTypes different = TestAllTypes.newBuilder().setSingleInt64(99L).build(); + ProtoLiteCelValueConverter distinctConverter = + ProtoLiteCelValueConverter.newInstance(DESCRIPTOR_POOL); + + new EqualsTester() + .addEqualityGroup( + ProtoMessageLiteValue.create( + TestAllTypes.getDefaultInstance(), + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER), + ProtoMessageLiteValue.create( + ByteString.EMPTY, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER)) + .addEqualityGroup( + ProtoMessageLiteValue.create( + populated, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER), + ProtoMessageLiteValue.create( + populated.toByteString(), + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER), + ProtoMessageLiteValue.create( + populated, "cel.expr.conformance.proto3.TestAllTypes", distinctConverter)) + .addEqualityGroup( + ProtoMessageLiteValue.create( + different, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER)) + .addEqualityGroup( + ProtoMessageLiteValue.create( + TestAllTypes.getDefaultInstance(), + "different.TypeName", + PROTO_LITE_CEL_VALUE_CONVERTER)) + .addEqualityGroup( + ProtoMessageLiteValue.create( + NestedMessage.getDefaultInstance(), + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", + PROTO_LITE_CEL_VALUE_CONVERTER)) + .testEquals(); + } + + @Test + public void create_withCorruptByteString_throwsOnValueMaterialization() { + ByteString corruptBytes = ByteString.copyFrom(new byte[] {(byte) 0xFF, (byte) 0xFF}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + corruptBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + + IllegalArgumentException thrown = + assertThrows(IllegalArgumentException.class, messageLiteValue::value); + + assertThat(thrown) + .hasMessageThat() + .contains( + "Failed to decode proto message of type: cel.expr.conformance.proto3.TestAllTypes"); + assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); + } + + @Test + public void create_withCorruptByteString_throwsOnSelect() { + ByteString corruptBytes = ByteString.copyFrom(new byte[] {(byte) 0xFF, (byte) 0xFF}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + corruptBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + + IllegalArgumentException thrown = + assertThrows(IllegalArgumentException.class, () -> messageLiteValue.select("single_int64")); + + assertThat(thrown) + .hasMessageThat() + .contains( + "Failed to decode proto message of type: cel.expr.conformance.proto3.TestAllTypes"); + assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); + } + + @Test + public void create_withCorruptByteString_throwsOnSelectByFieldNumber() { + ByteString corruptBytes = ByteString.copyFrom(new byte[] {0x10, (byte) 0x80}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + corruptBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + SelectField selectField = SelectField.create(2L, "single_int64", 3, 0L); + + IllegalArgumentException thrown = + assertThrows( + IllegalArgumentException.class, + () -> messageLiteValue.selectByFieldNumber(selectField)); + + assertThat(thrown) + .hasMessageThat() + .contains( + "Failed to decode proto message of type: cel.expr.conformance.proto3.TestAllTypes"); + assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); + } + + @Test + public void create_withCorruptByteString_throwsOnHasFieldByNumber() { + ByteString corruptBytes = ByteString.copyFrom(new byte[] {0x10, (byte) 0x80}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + corruptBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + SelectField selectField = SelectField.create(2L, "single_int64"); + + IllegalArgumentException thrown = + assertThrows( + IllegalArgumentException.class, () -> messageLiteValue.hasFieldByNumber(selectField)); + + assertThat(thrown) + .hasMessageThat() + .contains( + "Failed to decode proto message of type: cel.expr.conformance.proto3.TestAllTypes"); + assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); + } + @SuppressWarnings("ImmutableEnumChecker") // Test only private enum SelectFieldTestCase { BOOL("single_bool", true), @@ -125,6 +307,13 @@ private enum SelectFieldTestCase { REPEATED_DOUBLE("repeated_double", ImmutableList.of(3.5d, 4.5d)), REPEATED_STRING("repeated_string", ImmutableList.of("foo", "bar")), + REPEATED_NESTED_MESSAGE( + "repeated_nested_message", + ImmutableList.of( + ProtoMessageLiteValue.create( + NestedMessage.newBuilder().setBb(10).build(), + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", + PROTO_LITE_CEL_VALUE_CONVERTER))), MAP_INT64_INT64("map_int64_int64", ImmutableMap.of(1L, 2L, 3L, 4L)), @@ -193,6 +382,7 @@ public void selectField_success(@TestParameter SelectFieldTestCase testCase) { .addRepeatedDouble(4.5d) .addRepeatedString("foo") .addRepeatedString("bar") + .addRepeatedNestedMessage(NestedMessage.newBuilder().setBb(10)) .putMapStringString("a", "b") .putMapInt64Int64(1L, 2L) .putMapInt64Int64(3L, 4L) 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(); + } } diff --git a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java index c3334d7cb..7fdfe63bc 100644 --- a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java @@ -21,6 +21,7 @@ import com.google.common.collect.ImmutableCollection; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; import com.google.common.primitives.UnsignedLong; import com.google.protobuf.ByteString; import com.google.protobuf.CodedOutputStream; @@ -32,6 +33,8 @@ import dev.cel.common.internal.ProtoTimeUtils; import dev.cel.expr.conformance.proto3.NestedTestAllTypes; import dev.cel.expr.conformance.proto3.TestAllTypes; +import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage; +import dev.cel.expr.conformance.proto3.TestAllTypesCelDescriptor; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; import java.io.ByteArrayOutputStream; import java.io.IOException; @@ -62,10 +65,12 @@ private static Object decodeWireValue( @Test public void create_accessorsAndType() { ByteString bytes = ByteString.copyFromUtf8("test"); + RawProtoMessageLiteValue value = RawProtoMessageLiteValue.create(bytes, "custom.Message", EMPTY_CONVERTER); - assertThat(value.rawWireBytes()).isEqualTo(bytes); + assertThat(value.toByteString()).isEqualTo(bytes); + assertThat(value.protoTypeName()).isEqualTo("custom.Message"); assertThat(value.value()).isSameInstanceAs(value); assertThat(value.celType().name()).isEqualTo("custom.Message"); } @@ -76,7 +81,8 @@ public void create_defaultsUnknownMessageTypeName() { RawProtoMessageLiteValue value = RawProtoMessageLiteValue.create(bytes, EMPTY_CONVERTER); - assertThat(value.rawWireBytes()).isEqualTo(bytes); + assertThat(value.toByteString()).isEqualTo(bytes); + assertThat(value.protoTypeName()).isEqualTo("cel.@unknownMessage"); assertThat(value.celType().name()).isEqualTo("cel.@unknownMessage"); } @@ -86,10 +92,48 @@ public void create_emptyTypeName_normalizesToUnknownMessageTypeName() { RawProtoMessageLiteValue value = RawProtoMessageLiteValue.create(bytes, "", EMPTY_CONVERTER); - assertThat(value.rawWireBytes()).isEqualTo(bytes); + assertThat(value.toByteString()).isEqualTo(bytes); + assertThat(value.protoTypeName()).isEqualTo("cel.@unknownMessage"); assertThat(value.celType().name()).isEqualTo("cel.@unknownMessage"); } + @Test + @SuppressWarnings("SelfEquals") // Testing that equals throws even on self-comparison + public void equals_throwsUnsupportedOperationException() { + RawProtoMessageLiteValue value1 = + RawProtoMessageLiteValue.create(ByteString.EMPTY, "custom.Message", EMPTY_CONVERTER); + RawProtoMessageLiteValue value2 = + RawProtoMessageLiteValue.create(ByteString.EMPTY, "custom.Message", EMPTY_CONVERTER); + + UnsupportedOperationException selfThrown = + assertThrows(UnsupportedOperationException.class, () -> value1.equals(value1)); + UnsupportedOperationException otherThrown = + assertThrows(UnsupportedOperationException.class, () -> value1.equals(value2)); + + assertThat(selfThrown).hasMessageThat().isEqualTo("Message equality is not supported"); + assertThat(otherThrown).hasMessageThat().isEqualTo("Message equality is not supported"); + } + + @Test + public void hashCode_throwsUnsupportedOperationException() { + RawProtoMessageLiteValue value = + RawProtoMessageLiteValue.create(ByteString.EMPTY, "custom.Message", EMPTY_CONVERTER); + + UnsupportedOperationException thrown = + assertThrows(UnsupportedOperationException.class, value::hashCode); + + assertThat(thrown).hasMessageThat().isEqualTo("Message equality is not supported"); + } + + @Test + public void toString_returnsTypeNameAndByteSize() { + RawProtoMessageLiteValue value = + RawProtoMessageLiteValue.create( + ByteString.copyFromUtf8("secret"), "custom.Message", EMPTY_CONVERTER); + + assertThat(value.toString()).isEqualTo("WireMessageLite{protoTypeName=custom.Message, size=6}"); + } + @Test public void select_throwsCelAttributeNotFoundException() { RawProtoMessageLiteValue value = @@ -796,13 +840,14 @@ public void decodeWireEntries_repeatedMessage() { "sub.Message", /* isRepeated= */ true); - assertThat((Iterable) decoded) - .containsExactly( - RawProtoMessageLiteValue.create( - ByteString.copyFromUtf8("msg1"), "sub.Message", EMPTY_CONVERTER), - RawProtoMessageLiteValue.create( - ByteString.copyFromUtf8("msg2"), "sub.Message", EMPTY_CONVERTER)) - .inOrder(); + ImmutableList messages = (ImmutableList) decoded; + assertThat(messages).hasSize(2); + RawProtoMessageLiteValue msg0 = (RawProtoMessageLiteValue) messages.get(0); + RawProtoMessageLiteValue msg1 = (RawProtoMessageLiteValue) messages.get(1); + assertThat(msg0.toByteString()).isEqualTo(ByteString.copyFromUtf8("msg1")); + assertThat(msg0.protoTypeName()).isEqualTo("sub.Message"); + assertThat(msg1.toByteString()).isEqualTo(ByteString.copyFromUtf8("msg2")); + assertThat(msg1.protoTypeName()).isEqualTo("sub.Message"); } @Test @@ -980,13 +1025,13 @@ public void findByFieldNumber_intermediatePresent_returnsSubmessage() throws Exc Optional nav = raw.findByFieldNumber(SelectField.create(21L, "single_nested_message")); - RawProtoMessageLiteValue expected = - RawProtoMessageLiteValue.create( + Optional submessage = nav.map(RawProtoMessageLiteValue.class::cast); + assertThat(submessage.map(RawProtoMessageLiteValue::toByteString)) + .hasValue( ByteString.copyFrom(subBaos1.toByteArray()) - .concat(ByteString.copyFrom(subBaos2.toByteArray())), - "cel.@unknownMessage", - EMPTY_CONVERTER); - assertThat(nav).hasValue(expected); + .concat(ByteString.copyFrom(subBaos2.toByteArray()))); + assertThat(submessage.map(RawProtoMessageLiteValue::protoTypeName)) + .hasValue("cel.@unknownMessage"); } @Test @@ -1113,8 +1158,8 @@ public void selectByFieldNumber_absentMessageFieldWithoutDescriptor_returnsUnkno assertThat(selected).isInstanceOf(RawProtoMessageLiteValue.class); RawProtoMessageLiteValue message = (RawProtoMessageLiteValue) selected; - assertThat(message.rawWireBytes()).isEqualTo(ByteString.EMPTY); - assertThat(message.celType().name()).isEqualTo("cel.@unknownMessage"); + assertThat(message.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(message.protoTypeName()).isEqualTo("cel.@unknownMessage"); } @Test @@ -1307,14 +1352,12 @@ public void selectByFieldNumber_mapEntrySpecMissingMessageValue_returnsEmptyMess Object result = raw.selectByFieldNumber(newMapInt64NestedTypeField()); - assertThat(result) - .isEqualTo( - ImmutableMap.of( - 42L, - RawProtoMessageLiteValue.create( - ByteString.EMPTY, - "cel.expr.conformance.proto3.NestedTestAllTypes", - EMPTY_CONVERTER))); + ImmutableMap map = (ImmutableMap) result; + assertThat(map.keySet()).containsExactly(42L); + RawProtoMessageLiteValue messageValue = (RawProtoMessageLiteValue) map.get(42L); + assertThat(messageValue.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(messageValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } @Test @@ -1345,14 +1388,12 @@ public void selectByFieldNumber_mapEntrySpecRepeatedMessageValue_mergesFragments Object result = raw.selectByFieldNumber(newMapInt64NestedTypeField()); - assertThat(result) - .isEqualTo( - ImmutableMap.of( - 42L, - RawProtoMessageLiteValue.create( - fragment1.concat(fragment2), - "cel.expr.conformance.proto3.NestedTestAllTypes", - EMPTY_CONVERTER))); + ImmutableMap map = (ImmutableMap) result; + assertThat(map.keySet()).containsExactly(42L); + RawProtoMessageLiteValue messageValue = (RawProtoMessageLiteValue) map.get(42L); + assertThat(messageValue.toByteString()).isEqualTo(fragment1.concat(fragment2)); + assertThat(messageValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } @Test @@ -1804,4 +1845,57 @@ public void selectByFieldNumber_wireSubmessageWithProtoTypeName_decodesWithTypeN assertThat(((RawProtoMessageLiteValue) result).celType().name()) .isEqualTo("test.CustomMessage"); } + + @Test + public void selectByFieldNumber_unsetRegisteredSubmessage_returnsProtoMessageLiteValue() { + ProtoLiteCelValueConverter registeredConverter = + ProtoLiteCelValueConverter.newInstance( + DefaultLiteDescriptorPool.newInstance( + ImmutableSet.of(TestAllTypesCelDescriptor.getDescriptor()))); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + ByteString.EMPTY, "test.UnknownParent", registeredConverter); + SelectField field = + SelectField.create( + 999L, + "nested_msg", + FieldLiteDescriptor.Type.MESSAGE.getNumber(), + null, + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isInstanceOf(ProtoMessageLiteValue.class); + assertThat(((ProtoMessageLiteValue) result).value()) + .isEqualTo(NestedMessage.getDefaultInstance()); + } + + @Test + public void selectByFieldNumber_wireRegisteredSubmessage_decodesToProtoMessageLiteValue() + throws Exception { + ProtoLiteCelValueConverter registeredConverter = + ProtoLiteCelValueConverter.newInstance( + DefaultLiteDescriptorPool.newInstance( + ImmutableSet.of(TestAllTypesCelDescriptor.getDescriptor()))); + NestedMessage expectedNested = NestedMessage.newBuilder().setBb(42).build(); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + CodedOutputStream cos = CodedOutputStream.newInstance(baos); + cos.writeMessage(999, expectedNested); + cos.flush(); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + ByteString.copyFrom(baos.toByteArray()), "test.UnknownParent", registeredConverter); + SelectField field = + SelectField.create( + 999L, + "nested_msg", + FieldLiteDescriptor.Type.MESSAGE.getNumber(), + null, + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isInstanceOf(ProtoMessageLiteValue.class); + assertThat(((ProtoMessageLiteValue) result).value()).isEqualTo(expectedNested); + } } diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java index afe7efe02..c2d4631cf 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java @@ -44,7 +44,7 @@ import dev.cel.common.types.StructTypeReference; import dev.cel.common.values.CelByteString; import dev.cel.common.values.ProtoMessageLiteValueProvider; -import dev.cel.common.values.RawProtoMessageLiteValue; +import dev.cel.common.values.WireMessageLite; import dev.cel.expr.conformance.proto3.NestedTestAllTypes; import dev.cel.expr.conformance.proto3.NestedTestAllTypesCelDescriptor; import dev.cel.expr.conformance.proto3.TestAllTypes; @@ -535,15 +535,86 @@ public void select_populatedUnknownRepeatedScalar_decodesFromWireBytes( } @Test - public void select_populatedUnknownSubmessageLeaf_returnsRawMessage() throws Exception { - TestAllTypes msg = - TestAllTypes.newBuilder() - .setSingleNestedMessage(NestedMessage.newBuilder().setBb(123).build()) - .build(); + public void select_populatedUnknownSubmessageLeaf_withKnownSubmessageType_returnsMessageLite() + throws Exception { + NestedMessage nested = NestedMessage.newBuilder().setBb(123).build(); + TestAllTypes msg = TestAllTypes.newBuilder().setSingleNestedMessage(nested).build(); + + Object result = eval("msg.single_nested_message", msg); + + assertThat(result).isEqualTo(nested); + } + + @Test + public void select_unsetUnknownSubmessageLeaf_withKnownSubmessageType_returnsDefaultMessageLite() + throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); Object result = eval("msg.single_nested_message", msg); - assertThat(result).isInstanceOf(RawProtoMessageLiteValue.class); + assertThat(result).isEqualTo(NestedMessage.getDefaultInstance()); + } + + @Test + public void + select_populatedUnknownSubmessageLeaf_withoutSubmessageDescriptor_returnsWireMessageLite() + throws Exception { + NestedMessage nested = NestedMessage.newBuilder().setBb(123).build(); + TestAllTypes msg = TestAllTypes.newBuilder().setSingleNestedMessage(nested).build(); + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile("msg.single_nested_message").getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(result).isInstanceOf(WireMessageLite.class); + WireMessageLite wireMsg = (WireMessageLite) result; + assertThat(wireMsg.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + assertThat(wireMsg.toByteString()).isEqualTo(nested.toByteString()); + } + + @Test + public void + select_unsetUnknownSubmessageLeaf_withoutSubmessageDescriptor_returnsEmptyWireMessageLite() + throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile("msg.single_nested_message").getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(result).isInstanceOf(WireMessageLite.class); + WireMessageLite wireMsg = (WireMessageLite) result; + assertThat(wireMsg.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + assertThat(wireMsg.toByteString()).isEqualTo(ByteString.EMPTY); + } + + @Test + public void + select_populatedUnknownRepeatedSubmessageLeaf_withKnownSubmessageType_returnsMessageLiteList() + throws Exception { + Object result = eval("msg.repeated_nested_message", POPULATED_SERVER_MESSAGE); + + assertThat(result) + .isEqualTo( + ImmutableList.of( + NestedMessage.newBuilder().setBb(10).build(), + NestedMessage.newBuilder().setBb(20).build())); + } + + @Test + public void + select_populatedUnknownMapSubmessageLeaf_withKnownSubmessageType_returnsMessageLiteMap() + throws Exception { + Object result = eval("msg.map_string_message", POPULATED_SERVER_MESSAGE); + + assertThat(result) + .isEqualTo(ImmutableMap.of("m1", NestedMessage.newBuilder().setBb(55).build())); } @Test @@ -1111,69 +1182,49 @@ public void select_negativeInt32InPopulatedMessage_evaluatesCorrectly() throws E } @Test - public void submessageEquality_populatedSelfComparison_evaluatesTrue() throws Exception { - Object result = - eval("msg.single_nested_message == msg.single_nested_message", POPULATED_SERVER_MESSAGE); - - assertThat(result).isEqualTo(true); - } - - @Test - public void submessageEquality_unsetSelfComparison_evaluatesTrue() throws Exception { - Object result = - eval( + public void submessageEquality_withKnownSubmessageDescriptor_throwsUnsupportedOperationException( + @TestParameter({ "msg.single_nested_message == msg.single_nested_message", - TestAllTypes.getDefaultInstance()); - - assertThat(result).isEqualTo(true); - } - - @Test - public void submessageEquality_identicalSubmessagesAcrossDistinctPaths_evaluatesTrue() + "msg.repeated_nested_message == msg.repeated_nested_message" + }) + String expr) throws Exception { - NestedTestAllTypes nestedMsg = - NestedTestAllTypes.newBuilder() - .setPayload(POPULATED_SERVER_MESSAGE) - .setChild(NestedTestAllTypes.newBuilder().setPayload(POPULATED_SERVER_MESSAGE)) - .build(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile(expr).getAst()); + Program program = clientRuntime.createProgram(optimizedAst); - Object result = - evalNested( - "nested_msg.payload.single_nested_message ==" - + " nested_msg.child.payload.single_nested_message", - nestedMsg); + CelEvaluationException thrown = + assertThrows( + CelEvaluationException.class, + () -> program.eval(ImmutableMap.of("msg", POPULATED_SERVER_MESSAGE))); - assertThat(result).isEqualTo(true); + assertThat(thrown).hasCauseThat().isInstanceOf(UnsupportedOperationException.class); + assertThat(thrown).hasCauseThat().hasMessageThat().contains("Not implemented yet"); } @Test - public void submessageEquality_differentSubmessagesWithSameCelType_evaluatesFalse() + public void submessageEquality_withoutSubmessageDescriptor_throwsUnsupportedOperationException( + @TestParameter({ + "msg.single_nested_message == msg.single_nested_message", + "msg.repeated_nested_message == msg.repeated_nested_message" + }) + String expr) throws Exception { - Object differentResult = - eval( - "msg.repeated_nested_message[0] == msg.repeated_nested_message[1]", - POPULATED_SERVER_MESSAGE); - - assertThat(differentResult).isEqualTo(false); - } - - @Test - public void - submessageEquality_singularAndRepeatedUnknownSubmessagesWithIdenticalContent_evaluatesTrue() - throws Exception { - TestAllTypes msg = - TestAllTypes.newBuilder() - .setSingleNestedMessage(NestedMessage.newBuilder().setBb(123)) - .addRepeatedNestedMessage(NestedMessage.newBuilder().setBb(123)) - .build(); + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile(expr).getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); - Object result = - eval( - "msg.single_nested_message == msg.repeated_nested_message[0] &&" - + " !(msg.single_nested_message != msg.repeated_nested_message[0])", - msg); + CelEvaluationException thrown = + assertThrows( + CelEvaluationException.class, + () -> program.eval(ImmutableMap.of("msg", POPULATED_SERVER_MESSAGE))); - assertThat(result).isEqualTo(true); + assertThat(thrown).hasCauseThat().isInstanceOf(UnsupportedOperationException.class); + assertThat(thrown) + .hasCauseThat() + .hasMessageThat() + .contains("Message equality is not supported"); } @Test @@ -2248,6 +2299,109 @@ public void partialDescriptorRuntime_unoptimizedAst_throwsDescriptiveException( + " or enable SelectOptimizer."); } + @Test + public void descriptorlessRuntime_rootMessageResult_returnsWireMessageLite() throws Exception { + CelLiteRuntime descriptorlessRuntime = newDescriptorlessRuntime(); + CelAbstractSyntaxTree ast = serverCompiler.compile("msg").getAst(); + Program program = descriptorlessRuntime.createProgram(ast); + + Object result = program.eval(ImmutableMap.of("msg", POPULATED_SERVER_MESSAGE)); + + assertThat(result).isInstanceOf(WireMessageLite.class); + WireMessageLite wireMsg = (WireMessageLite) result; + assertThat(wireMsg.protoTypeName()).isEqualTo("cel.@unknownMessage"); + assertThat(wireMsg.toByteString()).isEqualTo(POPULATED_SERVER_MESSAGE.toByteString()); + } + + @Test + public void olderClassRoundTrip_descriptorlessRuntime_evaluatesFromReserializedUnknownFields( + @TestParameter DescriptorlessEvaluationTestCase testCase) throws Exception { + // NestedMessage only defines field 1 (int32 bb = 1); parsing a V2 TestAllTypes payload into it + // populates its known field (tag 1) alongside unknown V2 fields (tags 2..402) before + // re-serialization. + TestAllTypes v2Message = + testCase.message.equals(TestAllTypes.getDefaultInstance()) + ? testCase.message + : testCase.message.toBuilder().setSingleInt32(42).build(); + NestedMessage olderClassMsg = + NestedMessage.parseFrom(v2Message.toByteString(), ExtensionRegistryLite.getEmptyRegistry()); + CelLiteRuntime descriptorlessRuntime = newDescriptorlessRuntime(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile(testCase.expression).getAst()); + Program program = descriptorlessRuntime.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", olderClassMsg)); + + assertThat(result).isEqualTo(testCase.expectedResult); + } + + @Test + public void + olderClassRoundTrip_v1ClassDescriptor_evaluatesKnownAndUnknownFieldsAcrossSubmessages() + throws Exception { + TestAllTypes v2Payload = POPULATED_SERVER_MESSAGE.toBuilder().setSingleInt32(42).build(); + // Parse V2 TestAllTypes wire bytes into NestedMessage (which only knows field 1: int32 bb = 1). + // Field 1 is stored in NestedMessage's generated Java field, while all V2 fields (tags 2..402) + // are retained in NestedMessage's unknown fields. + NestedMessage v1RootMsg = + NestedMessage.parseFrom(v2Payload.toByteString(), ExtensionRegistryLite.getEmptyRegistry()); + NestedTestAllTypes v2NestedMsg = NestedTestAllTypes.newBuilder().setPayload(v2Payload).build(); + CelLiteRuntime v1OlderClassRuntime = newOlderClassV1Runtime(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize( + serverCompiler + .compile( + "has(msg.single_int32) && msg.single_int32 == 42" + + " && has(msg.single_int64) && msg.single_int64 == -42" + + " && msg.single_string == 'cel-skew-test'" + + " && msg.repeated_int64 == [10, 20]" + + " && msg.map_int32_int32[1] == 2" + + " && msg.single_duration == duration('1h')" + + " && has(nested_msg.payload.single_int32)" + + " && nested_msg.payload.single_int32 == 42" + + " && has(nested_msg.payload.single_int64)" + + " && nested_msg.payload.single_int64 == -42" + + " && nested_msg.payload.single_nested_message.bb == 123" + + " && nested_msg.payload.repeated_string == ['foo', 'bar']" + + " && nested_msg.payload.map_string_duration['d'] == duration('5m')") + .getAst()); + Program program = v1OlderClassRuntime.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", v1RootMsg, "nested_msg", v2NestedMsg)); + + assertThat(result).isEqualTo(true); + } + + @Test + public void + olderClassRoundTrip_v1ClassDescriptor_evaluatesUnsetDefaultsAndDecodesLeafSubmessageIntoOlderClass() + throws Exception { + TestAllTypes v2Payload = POPULATED_SERVER_MESSAGE.toBuilder().setSingleInt32(42).build(); + NestedMessage expectedV1Msg = + NestedMessage.parseFrom(v2Payload.toByteString(), ExtensionRegistryLite.getEmptyRegistry()); + NestedTestAllTypes v2NestedMsg = NestedTestAllTypes.newBuilder().setPayload(v2Payload).build(); + CelLiteRuntime v1OlderClassRuntime = newOlderClassV1Runtime(); + Program defaultCheckProgram = + v1OlderClassRuntime.createProgram( + serverOptimizer.optimize( + serverCompiler + .compile( + "!has(msg.single_int32) && msg.single_int32 == 0" + + " && !has(msg.single_int64) && msg.single_int64 == 0") + .getAst())); + Program leafSubmessageProgram = + v1OlderClassRuntime.createProgram( + serverOptimizer.optimize(serverCompiler.compile("nested_msg.payload").getAst())); + + Object defaultCheckResult = + defaultCheckProgram.eval(ImmutableMap.of("msg", NestedMessage.getDefaultInstance())); + Object leafSubmessageResult = + leafSubmessageProgram.eval(ImmutableMap.of("nested_msg", v2NestedMsg)); + + assertThat(defaultCheckResult).isEqualTo(true); + assertThat(leafSubmessageResult).isEqualTo(expectedV1Msg); + } + private Program compileScoreModelLateBoundProgram() throws Exception { Cel celWithLateFunc = serverCompiler @@ -2439,6 +2593,32 @@ private static CelLiteRuntime newDescriptorlessRuntime() { .build(); } + private static CelLiteRuntime newOlderClassV1Runtime() { + FieldLiteDescriptor v1SingleInt32Field = + new FieldLiteDescriptor( + /* fieldNumber= */ 1, + /* fieldName= */ "v1_single_int32", + /* javaType= */ FieldLiteDescriptor.JavaType.INT, + /* encodingType= */ FieldLiteDescriptor.EncodingType.SINGULAR, + /* protoFieldType= */ FieldLiteDescriptor.Type.INT32, + /* isPacked= */ false, + /* fieldProtoTypeName= */ ""); + MessageLiteDescriptor v1OlderClassTestAllTypesDesc = + new MessageLiteDescriptor( + TestAllTypes.getDescriptor().getFullName(), + ImmutableList.of(v1SingleInt32Field), + NestedMessage::newBuilder); + CelLiteDescriptor v1OlderClassDescriptor = + new CelLiteDescriptor("v1_older_class", ImmutableList.of(v1OlderClassTestAllTypesDesc)) {}; + return CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .setStandardFunctions(CelStandardFunctions.ALL_STANDARD_FUNCTIONS) + .setValueProvider( + ProtoMessageLiteValueProvider.newInstance( + v1OlderClassDescriptor, NestedTestAllTypesCelDescriptor.getDescriptor())) + .setContainer(CEL_CONTAINER) + .build(); + } + private CelOptimizer newSubexpressionOptimizer() { return CelOptimizerFactory.standardCelOptimizerBuilder(serverCompiler) .addAstOptimizers( diff --git a/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java b/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java index 50c39ed58..e67c4f124 100644 --- a/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java +++ b/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java @@ -22,6 +22,7 @@ import com.google.common.primitives.UnsignedLong; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.common.CelOptions; +import dev.cel.common.values.ProtoMessageLiteValueProvider; import dev.cel.expr.conformance.proto2.TestAllTypes; import org.junit.Test; import org.junit.runner.RunWith; @@ -60,13 +61,43 @@ public void objectEquals_messageLite_throws() { TestAllTypes.Builder builder = TestAllTypes.newBuilder(); TestAllTypes defaultInstance = TestAllTypes.getDefaultInstance(); - // Unimplemented until CelLiteDescriptor is available. - UnsupportedOperationException e = + UnsupportedOperationException builderThrown = assertThrows( UnsupportedOperationException.class, () -> runtimeEquality.objectEquals(builder, defaultInstance)); + UnsupportedOperationException stringThrown = + assertThrows( + UnsupportedOperationException.class, + () -> runtimeEquality.objectEquals("not_a_message", defaultInstance)); + + assertThat(builderThrown).hasMessageThat().isEqualTo("Not implemented yet"); + assertThat(stringThrown).hasMessageThat().isEqualTo("Not implemented yet"); + } + + @Test + public void objectEquals_wireMessageLite_throws() { + RuntimeEquality runtimeEquality = + RuntimeEquality.create(RuntimeHelpers.create(), CelOptions.DEFAULT); + Object wireMsg1 = + ProtoMessageLiteValueProvider.newInstance() + .protoCelValueConverter() + .toRuntimeValue(TestAllTypes.getDefaultInstance()); + Object wireMsg2 = + ProtoMessageLiteValueProvider.newInstance() + .protoCelValueConverter() + .toRuntimeValue(TestAllTypes.newBuilder().setSingleInt32(1).build()); + + UnsupportedOperationException messagesThrown = + assertThrows( + UnsupportedOperationException.class, + () -> runtimeEquality.objectEquals(wireMsg1, wireMsg2)); + UnsupportedOperationException stringThrown = + assertThrows( + UnsupportedOperationException.class, + () -> runtimeEquality.objectEquals("not_a_message", wireMsg1)); - assertThat(e).hasMessageThat().contains("Not implemented yet"); + assertThat(messagesThrown).hasMessageThat().isEqualTo("Message equality is not supported"); + assertThat(stringThrown).hasMessageThat().isEqualTo("Message equality is not supported"); } @Test