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 19e0db963..2a3a810a7 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java +++ b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java @@ -46,6 +46,7 @@ import java.util.Map; import java.util.Optional; import java.util.TreeMap; +import org.jspecify.annotations.Nullable; /** * {@code ProtoLiteCelValueConverter} handles bidirectional conversion between native Java and @@ -60,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; @@ -134,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: @@ -157,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)); @@ -172,6 +202,28 @@ Optional findFieldDescriptor(String protoTypeName, int fiel .flatMap(desc -> desc.findByFieldNumber(fieldNumber)); } + MessageLite parseMessageLite(ByteString bytes, String protoTypeName) { + MessageLiteDescriptor descriptor = descriptorPool.getDescriptorOrThrow(protoTypeName); + return parseMessageLite(bytes, descriptor); + } + + private static MessageLite parseMessageLite(ByteString bytes, MessageLiteDescriptor descriptor) { + if (bytes.isEmpty()) { + return descriptor.newMessageBuilder().getDefaultInstanceForType(); + } + return mergeMessageLite(descriptor.newMessageBuilder(), bytes, descriptor.getProtoTypeName()); + } + + private static MessageLite mergeMessageLite( + MessageLite.Builder builder, ByteString bytes, String protoTypeName) { + try { + return builder.mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()).build(); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + protoTypeName, e); + } + } + Optional tryDecodeProtoMessage(ByteString bytes, String protoTypeName) { return descriptorPool .findDescriptor(protoTypeName) @@ -180,14 +232,11 @@ Optional tryDecodeProtoMessage(ByteString bytes, String protoTypeName) { private Object decodeProtoMessage( ByteString bytes, String protoTypeName, MessageLiteDescriptor descriptor) { - try { - MessageLite.Builder builder = - descriptor.newMessageBuilder().mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()); - return toRuntimeValue(builder.build(), descriptor); - } catch (IOException e) { - throw new IllegalArgumentException( - "Failed to decode proto message of type: " + protoTypeName, e); + WellKnownProto wellKnownProto = WellKnownProto.getByTypeName(protoTypeName).orElse(null); + if (isStructLike(wellKnownProto)) { + return ProtoMessageLiteValue.create(bytes, protoTypeName, this); } + return fromWellKnownProto(parseMessageLite(bytes, descriptor), checkNotNull(wellKnownProto)); } @Override @@ -209,12 +258,11 @@ public Object toRuntimeValue(Object value) { private Object toRuntimeValue(MessageLite msg, MessageLiteDescriptor descriptor) { WellKnownProto wellKnownProto = WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null); - - if (wellKnownProto == null || wellKnownProto == WellKnownProto.FIELD_MASK) { + if (isStructLike(wellKnownProto)) { return ProtoMessageLiteValue.create(msg, descriptor.getProtoTypeName(), this); } - return fromWellKnownProto(msg, wellKnownProto); + return fromWellKnownProto(msg, checkNotNull(wellKnownProto)); } private Object getDefaultValue(FieldLiteDescriptor fieldDescriptor) { @@ -260,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); } @@ -286,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); @@ -299,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: @@ -347,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; @@ -371,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 { @@ -421,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 f1a738d74..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: * *

    @@ -55,24 +62,45 @@ 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()); @@ -148,7 +196,24 @@ static ProtoMessageLiteValue create( 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/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java index 5337a2df0..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; @@ -49,6 +50,7 @@ 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; @@ -58,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 @@ -176,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)); } @@ -260,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()) @@ -277,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(); @@ -318,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( @@ -410,7 +412,7 @@ public void readAllFields_nestedMessageWithoutDescriptor_returnsRawProtoMessageL MessageFields fields = PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( - msg.toByteArray(), "cel.expr.conformance.proto3.TestAllTypes"); + msg.toByteString(), "cel.expr.conformance.proto3.TestAllTypes"); assertThat(fields.values().keySet()).containsExactly("oneof_type"); Object fieldValue = fields.values().get("oneof_type"); @@ -433,13 +435,16 @@ public void tryDecodeProtoMessage_wellKnownType_returnsDecodedValue() { } @Test - public void tryDecodeProtoMessage_registeredMessageType_returnsProtoMessageLiteValue() { + public void tryDecodeProtoMessage_registeredMessageType_returnsWireBackedProtoMessageLiteValue() { NestedMessage nestedMsg = NestedMessage.newBuilder().setBb(42).build(); Optional decoded = 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( @@ -448,6 +453,19 @@ public void tryDecodeProtoMessage_registeredMessageType_returnsProtoMessageLiteV PROTO_LITE_CEL_VALUE_CONVERTER)); } + @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 tryDecodeProtoMessage_registeredMessageTypeEmptyBytes_returnsDefaultProtoMessageLiteValue() { @@ -502,4 +520,71 @@ public void tryDecodeProtoMessage_anyType_throwsUnsupportedOperationException() 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)