From 2ad288606f71da03bde58a9ce9e7608773eb6868 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Mon, 5 Oct 2026 11:13:56 -0700 Subject: [PATCH] Pass over CodedInputStream once and lazily slice ByteString submessages PiperOrigin-RevId: 993798210 --- .../java/dev/cel/common/values/BUILD.bazel | 2 + .../values/ProtoLiteCelValueConverter.java | 220 +++++++------ .../common/values/ProtoMessageLiteValue.java | 89 +++++- .../values/RawProtoMessageLiteValue.java | 74 +++-- .../cel/common/values/WireMessageLite.java | 47 +++ .../ProtoLiteCelValueConverterTest.java | 180 ++++++++--- .../values/ProtoMessageLiteValueTest.java | 147 +++++++++ .../values/RawProtoMessageLiteValueTest.java | 162 ++++++++-- .../CelLiteRuntimeVersionSkewTest.java | 300 ++++++++++++++---- .../dev/cel/runtime/RuntimeEqualityTest.java | 37 ++- 10 files changed, 992 insertions(+), 266 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..785a2ef32 100644 --- a/common/src/main/java/dev/cel/common/values/BUILD.bazel +++ b/common/src/main/java/dev/cel/common/values/BUILD.bazel @@ -320,6 +320,7 @@ java_library( "ProtoLiteCelValueConverter.java", "ProtoMessageLiteValue.java", "RawProtoMessageLiteValue.java", + "WireMessageLite.java", ], tags = [ ], @@ -350,6 +351,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..2a3a810a7 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; @@ -45,9 +44,9 @@ 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 +61,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 +135,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 +154,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 +202,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 +249,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,12 +308,7 @@ 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); } @@ -312,7 +329,7 @@ private Map.Entry readSingleMapEntry( CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException { String entryTypeName = fieldDescriptor.getFieldProtoTypeName(); ImmutableMap singleMapEntry = - readAllFields(inputStream.readByteArray(), entryTypeName).values(); + readAllFields(inputStream.readBytes(), entryTypeName).values(); Object key = singleMapEntry.get(MAP_KEY_FIELD_NAME); if (key == null) { key = getDefaultCelValue(entryTypeName, MAP_KEY_FIELD_NAME); @@ -325,25 +342,32 @@ private Map.Entry readSingleMapEntry( return new AbstractMap.SimpleEntry<>(key, value); } - 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; } + String fieldName = fieldDescriptor.getFieldName(); Object payload; switch (tagWireType) { case WireFormat.WIRETYPE_VARINT: @@ -373,18 +397,24 @@ MessageFields readAllFields(byte[] bytes, String protoTypeName) throws IOExcepti + protoFieldType); } - payload = readLengthDelimitedField(inputStream, fieldDescriptor); + payload = + readLengthDelimitedField( + inputStream, fieldDescriptor, /* existingValue= */ null); } break; case MAP: + // Safe because MAP fields only ever store a LinkedHashMap in fieldValues. + @SuppressWarnings("unchecked") Map fieldMap = - mapFieldValues.computeIfAbsent(fieldNumber, (unused) -> new LinkedHashMap<>()); + (Map) + fieldValues.computeIfAbsent(fieldName, (unused) -> new LinkedHashMap<>()); Map.Entry mapEntry = readSingleMapEntry(inputStream, fieldDescriptor); fieldMap.put(mapEntry.getKey(), mapEntry.getValue()); - payload = fieldMap; - break; + continue; default: - payload = readLengthDelimitedField(inputStream, fieldDescriptor); + payload = + readLengthDelimitedField( + inputStream, fieldDescriptor, fieldValues.get(fieldName)); break; } break; @@ -397,30 +427,27 @@ MessageFields readAllFields(byte[] bytes, String protoTypeName) throws IOExcepti } if (fieldDescriptor.getEncodingType().equals(EncodingType.LIST)) { - String fieldName = fieldDescriptor.getFieldName(); - List repeatedValues = - repeatedFieldValues.computeIfAbsent(fieldNumber, (unused) -> new ArrayList<>()); - if (payload instanceof Collection) { - repeatedValues.addAll((Collection) payload); + Collection elements = (Collection) payload; + if (!elements.isEmpty()) { + getOrCreateRepeatedList(fieldValues, fieldName).addAll(elements); + } } else { - repeatedValues.add(payload); - } - if (!repeatedValues.isEmpty()) { - fieldValues.put(fieldName, repeatedValues); + getOrCreateRepeatedList(fieldValues, fieldName).add(payload); } } else { - fieldValues.put(fieldDescriptor.getFieldName(), payload); + fieldValues.put(fieldName, payload); } } - // 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); + return MessageFields.create(ImmutableMap.copyOf(fieldValues), unknownFields); } - MessageFields readMessageFields(MessageLite msg, String protoTypeName) throws IOException { - return readAllFields(msg.toByteArray(), protoTypeName); + // Safe because LIST fields only ever store an ArrayList in fieldValues. + @SuppressWarnings("unchecked") + private static List getOrCreateRepeatedList( + Map fieldValues, String fieldName) { + return (List) fieldValues.computeIfAbsent(fieldName, (unused) -> new ArrayList<>()); } static Object readUnknownField(int tagWireType, CodedInputStream inputStream) throws IOException { @@ -447,19 +474,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..175392e5f 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: * *

    @@ -52,27 +59,48 @@ */ @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()); + } + + ByteString toByteString() { + ByteString bytes = wireBytes(); + return bytes != null ? bytes : checkNotNull(rawValue()).toByteString(); + } + @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 +110,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,9 +142,8 @@ 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 @@ -130,7 +178,7 @@ public Optional findByFieldNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { return Optional.ofNullable(fieldValues().get(fd.getFieldName())) - .map(value -> protoLiteCelValueConverter().toRuntimeValue(value)); + .map(protoLiteCelValueConverter()::toRuntimeValue); } return RawProtoMessageLiteValue.navigateWire( field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); @@ -142,13 +190,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/RawProtoMessageLiteValue.java b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java index b0c41f058..65ace1c7b 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(); } /** @@ -169,27 +194,15 @@ 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); } /** @@ -419,9 +432,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 +448,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 +527,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..cc933b5ff 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java @@ -25,6 +25,7 @@ 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; @@ -45,9 +46,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 +60,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 +104,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 +178,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 +262,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 +279,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 +320,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 +392,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 +412,179 @@ 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); + } } 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..822de79b3 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,143 @@ 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); + } + @SuppressWarnings("ImmutableEnumChecker") // Test only private enum SelectFieldTestCase { BOOL("single_bool", true), @@ -125,6 +264,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 +339,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/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