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 2a3a810a7..bdcb4cd6d 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java +++ b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java @@ -40,7 +40,6 @@ import java.io.IOException; import java.util.AbstractMap; import java.util.ArrayList; -import java.util.Collection; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; @@ -313,33 +312,90 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) { throw new IllegalStateException("Unexpected java type: " + type); } - private ImmutableList readPackedRepeatedFields( + private Map.Entry readSingleMapEntry( CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException { + String entryTypeName = fieldDescriptor.getFieldProtoTypeName(); + MessageLiteDescriptor entryDescriptor = descriptorPool.getDescriptorOrThrow(entryTypeName); + FieldLiteDescriptor keyDescriptor = entryDescriptor.getByFieldNameOrThrow(MAP_KEY_FIELD_NAME); + FieldLiteDescriptor valueDescriptor = + entryDescriptor.getByFieldNameOrThrow(MAP_VALUE_FIELD_NAME); int length = inputStream.readInt32(); int oldLimit = inputStream.pushLimit(length); - ImmutableList.Builder builder = ImmutableList.builder(); + Object key = null; + Object value = null; while (inputStream.getBytesUntilLimit() > 0) { - builder.add(readPrimitiveField(inputStream, fieldDescriptor)); + int tag = inputStream.readTag(); + int tagWireType = WireFormat.getTagWireType(tag); + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber == keyDescriptor.getFieldNumber()) { + key = readSingularField(tagWireType, inputStream, keyDescriptor, key); + } else if (fieldNumber == valueDescriptor.getFieldNumber()) { + value = readSingularField(tagWireType, inputStream, valueDescriptor, value); + } else { + skipWireField(tag, inputStream); + } } inputStream.popLimit(oldLimit); - return builder.build(); - } - - private Map.Entry readSingleMapEntry( - CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException { - String entryTypeName = fieldDescriptor.getFieldProtoTypeName(); - ImmutableMap singleMapEntry = - readAllFields(inputStream.readBytes(), entryTypeName).values(); - Object key = singleMapEntry.get(MAP_KEY_FIELD_NAME); if (key == null) { - key = getDefaultCelValue(entryTypeName, MAP_KEY_FIELD_NAME); + key = getDefaultCelValue(keyDescriptor); } - Object value = singleMapEntry.get(MAP_VALUE_FIELD_NAME); if (value == null) { - value = getDefaultCelValue(entryTypeName, MAP_VALUE_FIELD_NAME); + value = getDefaultCelValue(valueDescriptor); } - return new AbstractMap.SimpleEntry<>(key, value); + return new AbstractMap.SimpleImmutableEntry<>(key, value); + } + + @Nullable Object readSingleField(ByteString bytes, FieldLiteDescriptor fieldDescriptor) + throws IOException { + CodedInputStream inputStream = bytes.newCodedInput(); + int targetFieldNumber = fieldDescriptor.getFieldNumber(); + Object fieldValue = null; + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber != targetFieldNumber) { + skipWireField(tag, inputStream); + continue; + } + int tagWireType = WireFormat.getTagWireType(tag); + fieldValue = readFieldValue(tagWireType, inputStream, fieldDescriptor, fieldValue); + } + return fieldValue; + } + + boolean hasSingleField(ByteString bytes, FieldLiteDescriptor fieldDescriptor) throws IOException { + return hasSingleField( + bytes, + fieldDescriptor.getFieldNumber(), + fieldDescriptor.getEncodingType().equals(EncodingType.LIST) + && fieldDescriptor.getIsPacked()); + } + + static boolean hasSingleField(ByteString bytes, int targetFieldNumber, boolean isPackableList) + throws IOException { + CodedInputStream inputStream = bytes.newCodedInput(); + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber != targetFieldNumber) { + skipWireField(tag, inputStream); + continue; + } + int tagWireType = WireFormat.getTagWireType(tag); + // In protobuf wire format, a zero-length entry for a singular field (e.g. empty string, + // bytes, or empty submessage) represents explicit presence on the wire. Only packed + // repeated fields with empty payload represent an empty/absent collection. + if (isPackableList && tagWireType == WireFormat.WIRETYPE_LENGTH_DELIMITED) { + int length = inputStream.readInt32(); + inputStream.skipRawBytes(length); + if (length > 0) { + return true; + } + continue; + } + skipWireField(tag, inputStream); + return true; + } + return false; } MessageFields readAllFields(ByteString bytes, String protoTypeName) throws IOException { @@ -368,86 +424,131 @@ private MessageFields readAllFields( } String fieldName = fieldDescriptor.getFieldName(); - Object payload; - switch (tagWireType) { - case WireFormat.WIRETYPE_VARINT: - payload = readPrimitiveField(inputStream, fieldDescriptor); - break; - case WireFormat.WIRETYPE_FIXED32: - payload = readFixed32BitField(inputStream, fieldDescriptor); - break; - case WireFormat.WIRETYPE_FIXED64: - payload = readFixed64BitField(inputStream, fieldDescriptor); - break; - case WireFormat.WIRETYPE_LENGTH_DELIMITED: - EncodingType encodingType = fieldDescriptor.getEncodingType(); - switch (encodingType) { - case LIST: - if (fieldDescriptor.getIsPacked()) { - payload = readPackedRepeatedFields(inputStream, fieldDescriptor); - } else { - FieldLiteDescriptor.Type protoFieldType = fieldDescriptor.getProtoFieldType(); - boolean isLenDelimited = - protoFieldType.equals(FieldLiteDescriptor.Type.MESSAGE) - || protoFieldType.equals(FieldLiteDescriptor.Type.STRING) - || protoFieldType.equals(FieldLiteDescriptor.Type.BYTES); - if (!isLenDelimited) { - throw new IllegalStateException( - "Unexpected field type encountered for LEN-Delimited record: " - + protoFieldType); - } - - payload = - readLengthDelimitedField( - inputStream, fieldDescriptor, /* existingValue= */ null); - } - break; - case MAP: - // Safe because MAP fields only ever store a LinkedHashMap in fieldValues. - @SuppressWarnings("unchecked") - Map fieldMap = - (Map) - fieldValues.computeIfAbsent(fieldName, (unused) -> new LinkedHashMap<>()); - Map.Entry mapEntry = readSingleMapEntry(inputStream, fieldDescriptor); - fieldMap.put(mapEntry.getKey(), mapEntry.getValue()); - continue; - default: - payload = - readLengthDelimitedField( - inputStream, fieldDescriptor, fieldValues.get(fieldName)); - break; - } - break; - case WireFormat.WIRETYPE_START_GROUP: - case WireFormat.WIRETYPE_END_GROUP: - // TODO: Support groups - throw new UnsupportedOperationException("Groups are not supported"); - default: - throw new IllegalArgumentException("Unexpected wire type: " + tagWireType); - } - - if (fieldDescriptor.getEncodingType().equals(EncodingType.LIST)) { - if (payload instanceof Collection) { - Collection elements = (Collection) payload; - if (!elements.isEmpty()) { - getOrCreateRepeatedList(fieldValues, fieldName).addAll(elements); - } - } else { - getOrCreateRepeatedList(fieldValues, fieldName).add(payload); - } - } else { - fieldValues.put(fieldName, payload); + Object fieldValue = + readFieldValue(tagWireType, inputStream, fieldDescriptor, fieldValues.get(fieldName)); + if (fieldValue != null) { + fieldValues.put(fieldName, fieldValue); } } return MessageFields.create(ImmutableMap.copyOf(fieldValues), unknownFields); } - // Safe because LIST fields only ever store an ArrayList in fieldValues. + private @Nullable Object readFieldValue( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + EncodingType encodingType = fieldDescriptor.getEncodingType(); + switch (encodingType) { + case SINGULAR: + return readSingularField(tagWireType, inputStream, fieldDescriptor, existingValue); + case LIST: + return readRepeatedField(tagWireType, inputStream, fieldDescriptor, existingValue); + case MAP: + return readMapField(tagWireType, inputStream, fieldDescriptor, existingValue); + } + throw new IllegalStateException("Unexpected encoding type: " + encodingType); + } + + private Object readSingularField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + switch (tagWireType) { + case WireFormat.WIRETYPE_VARINT: + return readPrimitiveField(inputStream, fieldDescriptor); + case WireFormat.WIRETYPE_FIXED32: + return readFixed32BitField(inputStream, fieldDescriptor); + case WireFormat.WIRETYPE_FIXED64: + return readFixed64BitField(inputStream, fieldDescriptor); + case WireFormat.WIRETYPE_LENGTH_DELIMITED: + return readLengthDelimitedField(inputStream, fieldDescriptor, existingValue); + case WireFormat.WIRETYPE_START_GROUP: + case WireFormat.WIRETYPE_END_GROUP: + throw new UnsupportedOperationException("Groups are not supported"); + default: + throw new IllegalArgumentException("Unexpected wire type: " + tagWireType); + } + } + + // Safe because LIST fields only ever store an ArrayList as their accumulated value. + @SuppressWarnings("unchecked") + private @Nullable List readRepeatedField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + List repeatedValues = (List) existingValue; + if (tagWireType == WireFormat.WIRETYPE_LENGTH_DELIMITED && fieldDescriptor.getIsPacked()) { + return readPackedRepeatedFields(inputStream, fieldDescriptor, repeatedValues); + } + Object element = + readSingularField(tagWireType, inputStream, fieldDescriptor, /* existingValue= */ null); + if (repeatedValues == null) { + repeatedValues = new ArrayList<>(); + } + repeatedValues.add(element); + return repeatedValues; + } + + private static @Nullable List readPackedRepeatedFields( + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable List repeatedValues) + throws IOException { + int length = inputStream.readInt32(); + if (length == 0) { + return repeatedValues; + } + int oldLimit = inputStream.pushLimit(length); + if (repeatedValues == null) { + repeatedValues = new ArrayList<>(); + } + while (inputStream.getBytesUntilLimit() > 0) { + repeatedValues.add(readPrimitiveField(inputStream, fieldDescriptor)); + } + inputStream.popLimit(oldLimit); + return repeatedValues; + } + + // Safe because MAP fields only ever store a LinkedHashMap as their accumulated value. @SuppressWarnings("unchecked") - private static List getOrCreateRepeatedList( - Map fieldValues, String fieldName) { - return (List) fieldValues.computeIfAbsent(fieldName, (unused) -> new ArrayList<>()); + private Map readMapField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + if (tagWireType != WireFormat.WIRETYPE_LENGTH_DELIMITED) { + throw new IllegalStateException("Unexpected wire type for map field: " + tagWireType); + } + Map mapValues = + existingValue != null ? (Map) existingValue : new LinkedHashMap<>(); + Map.Entry mapEntry = readSingleMapEntry(inputStream, fieldDescriptor); + mapValues.put(mapEntry.getKey(), mapEntry.getValue()); + return mapValues; + } + + static void skipWireField(int tag, CodedInputStream inputStream) throws IOException { + int tagWireType = WireFormat.getTagWireType(tag); + switch (tagWireType) { + case WireFormat.WIRETYPE_VARINT: + case WireFormat.WIRETYPE_FIXED64: + case WireFormat.WIRETYPE_LENGTH_DELIMITED: + case WireFormat.WIRETYPE_FIXED32: + inputStream.skipField(tag); + return; + case WireFormat.WIRETYPE_START_GROUP: + case WireFormat.WIRETYPE_END_GROUP: + throw new UnsupportedOperationException("Groups are not supported"); + default: + throw new IllegalArgumentException("Unknown wire type: " + tagWireType); + } } static Object readUnknownField(int tagWireType, CodedInputStream inputStream) throws IOException { 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 175392e5f..b1fe106db 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java @@ -52,9 +52,9 @@ * resolving by {@link SelectField#fieldNumber()} maps the number to the runtime descriptor's * current field name, preventing {@code CelAttributeNotFoundException}. *
  • Version skew / unknown fields: When evaluating payloads serialized by a newer binary - * containing fields absent from the local {@code CelLiteDescriptor}, the unknown wire bytes - * are preserved in {@link #unknownFields()} and decoded on demand using the compile-time wire - * type and default metadata in {@link SelectField}. + * containing fields absent from the local {@code CelLiteDescriptor}, the unknown fields are + * decoded on demand directly from the message's wire bytes using the compile-time wire type + * and default metadata in {@link SelectField}. * */ @AutoValue @@ -85,9 +85,14 @@ public MessageLite value() { .parseMessageLite(checkNotNull(wireBytes()), celType().name()); } + @Memoized + ByteString serializedRawValue() { + return checkNotNull(rawValue()).toByteString(); + } + ByteString toByteString() { ByteString bytes = wireBytes(); - return bytes != null ? bytes : checkNotNull(rawValue()).toByteString(); + return bytes != null ? bytes : serializedRawValue(); } @Memoized @@ -150,7 +155,7 @@ public Optional find(String field) { public Object selectByFieldNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { - Object known = fieldValues().get(fd.getFieldName()); + Object known = readField(fd); if (known != null) { return protoLiteCelValueConverter().toRuntimeValue(known); } @@ -160,28 +165,48 @@ public Object selectByFieldNumber(SelectField field) { return protoLiteCelValueConverter().getDefaultCelValue(fd); } return RawProtoMessageLiteValue.selectWireOrDefault( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, + RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()), + protoLiteCelValueConverter()); } @Override public boolean hasFieldByNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { - return fieldValues().containsKey(fd.getFieldName()); + return hasField(fd); } - return RawProtoMessageLiteValue.isPresentInWire( - field, unknownFields().get(field.fieldNumber())); + return RawProtoMessageLiteValue.isPresentInWire(toByteString(), field); } @Override public Optional findByFieldNumber(SelectField field) { FieldLiteDescriptor fd = findFieldDescriptor(field); if (fd != null) { - return Optional.ofNullable(fieldValues().get(fd.getFieldName())) - .map(protoLiteCelValueConverter()::toRuntimeValue); + return Optional.ofNullable(readField(fd)).map(protoLiteCelValueConverter()::toRuntimeValue); } return RawProtoMessageLiteValue.navigateWire( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, + RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()), + protoLiteCelValueConverter()); + } + + private @Nullable Object readField(FieldLiteDescriptor fd) { + try { + return protoLiteCelValueConverter().readSingleField(toByteString(), fd); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } + } + + private boolean hasField(FieldLiteDescriptor fd) { + try { + return protoLiteCelValueConverter().hasSingleField(toByteString(), fd); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } } private @Nullable FieldLiteDescriptor findFieldDescriptor(SelectField field) { 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 65ace1c7b..e6bd6466b 100644 --- a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java @@ -153,18 +153,45 @@ private CelAttributeNotFoundException newUnoptimizedFieldResolutionException(Str @Override public Object selectByFieldNumber(SelectField field) { return selectWireOrDefault( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, readWireField(toByteString(), field.fieldNumber()), protoLiteCelValueConverter()); } @Override public boolean hasFieldByNumber(SelectField field) { - return isPresentInWire(field, unknownFields().get(field.fieldNumber())); + return isPresentInWire(toByteString(), field); } @Override public Optional findByFieldNumber(SelectField field) { return navigateWire( - field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter()); + field, readWireField(toByteString(), field.fieldNumber()), protoLiteCelValueConverter()); + } + + /** + * Scans {@code wireBytes} for a single {@code targetFieldNumber}, skipping all other wire tags. + * + *

    Package-private: shared with {@code ProtoMessageLiteValue} for unknown field resolution. + */ + static ImmutableList readWireField(ByteString wireBytes, int targetFieldNumber) { + if (wireBytes.isEmpty()) { + return ImmutableList.of(); + } + ImmutableList.Builder entries = ImmutableList.builder(); + try { + CodedInputStream inputStream = wireBytes.newCodedInput(); + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + int fieldNumber = WireFormat.getTagFieldNumber(tag); + if (fieldNumber != targetFieldNumber) { + ProtoLiteCelValueConverter.skipWireField(tag, inputStream); + continue; + } + int tagWireType = WireFormat.getTagWireType(tag); + entries.add(ProtoLiteCelValueConverter.readUnknownField(tagWireType, inputStream)); + } + } catch (IOException e) { + throw new IllegalStateException("Failed to parse raw proto message wire bytes", e); + } + return entries.build(); } /** @@ -206,31 +233,28 @@ private static Object resolveDefault(SelectField field, ProtoLiteCelValueConvert } /** - * Returns whether a field has presence in preserved wire bytes. + * Scans {@code wireBytes} to determine whether {@code field} is present on the wire. * *

    Package-private: shared with {@code ProtoMessageLiteValue} for unknown field resolution. */ - static boolean isPresentInWire(SelectField field, ImmutableList unknowns) { + static boolean isPresentInWire(ByteString wireBytes, SelectField field) { + try { + return ProtoLiteCelValueConverter.hasSingleField( + wireBytes, field.fieldNumber(), isPackableRepeated(field)); + } catch (IOException e) { + throw new IllegalStateException("Failed to parse raw proto message wire bytes", e); + } + } + + private static boolean isPresentInWire(SelectField field, ImmutableList unknowns) { if (unknowns.isEmpty()) { return false; } - boolean isRepeated = field.defaultValue() instanceof List; - int typeCode = field.typeCode(); - // In protobuf wire format, a zero-length entry for a singular field (e.g. empty string, // bytes, or empty submessage) represents explicit presence on the wire. Only packed repeated // fields with empty payload represent an empty/absent collection. - if (!isRepeated) { - return true; - } - - boolean isPackable = - typeCode != FieldLiteDescriptor.Type.STRING.getNumber() - && typeCode != FieldLiteDescriptor.Type.BYTES.getNumber() - && typeCode != FieldLiteDescriptor.Type.MESSAGE.getNumber() - && typeCode != FieldLiteDescriptor.Type.GROUP.getNumber(); - if (!isPackable) { + if (!isPackableRepeated(field)) { return true; } @@ -242,6 +266,13 @@ static boolean isPresentInWire(SelectField field, ImmutableList unknowns return false; } + private static boolean isPackableRepeated(SelectField field) { + return (field.defaultValue() instanceof List) + && FieldLiteDescriptor.Type.forNumber(field.typeCode()) + .toWireFormatFieldType() + .isPackable(); + } + /** * Navigates a field on preserved wire bytes, returning empty if absent. * 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 cc933b5ff..a0fab9ca9 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java @@ -19,6 +19,7 @@ import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableListMultimap; +import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.google.common.collect.Multimap; import com.google.common.primitives.UnsignedLong; @@ -39,6 +40,7 @@ import com.google.protobuf.Timestamp; import com.google.protobuf.UInt32Value; import com.google.protobuf.UInt64Value; +import com.google.protobuf.WireFormat; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.common.internal.CelLiteDescriptorPool; @@ -587,4 +589,185 @@ public void readAllFields_emptyBytes_returnsEmptySingleton() throws Exception { assertThat(fields).isSameInstanceAs(MessageFields.EMPTY); } + + @Test + public void readSingleField_emptyBytes_returnsNull() throws Exception { + FieldLiteDescriptor fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", TestAllTypes.SINGLE_INT64_FIELD_NUMBER) + .get(); + + Object result = PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(ByteString.EMPTY, fd); + + assertThat(result).isNull(); + } + + @Test + public void readSingleField_absentField_returnsNull() throws Exception { + TestAllTypes proto = TestAllTypes.newBuilder().setSingleInt64(42L).build(); + FieldLiteDescriptor fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", TestAllTypes.SINGLE_BOOL_FIELD_NUMBER) + .get(); + + Object result = PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(proto.toByteString(), fd); + + assertThat(result).isNull(); + } + + @Test + public void readSingleField_emptyPackedRepeated_returnsNull() throws Exception { + ByteArrayOutputStream emptyPackedOut = new ByteArrayOutputStream(); + CodedOutputStream emptyPackedCos = CodedOutputStream.newInstance(emptyPackedOut); + emptyPackedCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + emptyPackedCos.flush(); + ByteString emptyPackedBytes = ByteString.copyFrom(emptyPackedOut.toByteArray()); + FieldLiteDescriptor repeatedInt32Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", + TestAllTypes.REPEATED_INT32_FIELD_NUMBER) + .get(); + + Object result = + PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(emptyPackedBytes, repeatedInt32Fd); + + assertThat(result).isNull(); + } + + @SuppressWarnings("ImmutableEnumChecker") // Test only + private enum ReadSingleFieldTestCase { + SINGLE_STRING(TestAllTypes.SINGLE_STRING_FIELD_NUMBER, "target_str"), + REPEATED_STRING(TestAllTypes.REPEATED_STRING_FIELD_NUMBER, ImmutableList.of("a", "b")), + REPEATED_INT32(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, ImmutableList.of(1, 2)), + MAP_STRING_STRING( + TestAllTypes.MAP_STRING_STRING_FIELD_NUMBER, ImmutableMap.of("k", "v", "k2", "v2", "", "")), + SPLIT_DURATION( + TestAllTypes.SINGLE_DURATION_FIELD_NUMBER, + Duration.newBuilder().setSeconds(10).setNanos(500).build()); + + private final int fieldNumber; + private final Object expected; + + ReadSingleFieldTestCase(int fieldNumber, Object expected) { + this.fieldNumber = fieldNumber; + this.expected = expected; + } + } + + @Test + public void readSingleField_skipsOtherFieldsAndDecodesTarget( + @TestParameter ReadSingleFieldTestCase testCase) throws Exception { + TestAllTypes part1 = + TestAllTypes.newBuilder() + .setSingleInt64(42L) + .setSingleFixed32(10) + .setSingleFixed64(20L) + .setSingleString("target_str") + .addRepeatedString("a") + .addRepeatedInt32(1) + .putMapStringString("k", "v") + .setSingleDuration(Duration.newBuilder().setSeconds(10)) + .build(); + TestAllTypes part2 = + TestAllTypes.newBuilder() + .addRepeatedString("b") + .addRepeatedInt32(2) + .putMapStringString("k2", "v2") + .setSingleDuration(Duration.newBuilder().setNanos(500)) + .build(); + ByteArrayOutputStream mapEntryWithUnknownOut = new ByteArrayOutputStream(); + CodedOutputStream mapEntryWithUnknownCos = + CodedOutputStream.newInstance(mapEntryWithUnknownOut); + mapEntryWithUnknownCos.writeInt64(3, 99L); + mapEntryWithUnknownCos.flush(); + ByteArrayOutputStream extraWireOut = new ByteArrayOutputStream(); + CodedOutputStream extraWireCos = CodedOutputStream.newInstance(extraWireOut); + extraWireCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + extraWireCos.writeByteArray( + TestAllTypes.MAP_STRING_STRING_FIELD_NUMBER, mapEntryWithUnknownOut.toByteArray()); + extraWireCos.flush(); + ByteString bytes = + part1 + .toByteString() + .concat(part2.toByteString()) + .concat(ByteString.copyFrom(extraWireOut.toByteArray())); + FieldLiteDescriptor fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor("cel.expr.conformance.proto3.TestAllTypes", testCase.fieldNumber) + .get(); + + Object result = PROTO_LITE_CEL_VALUE_CONVERTER.readSingleField(bytes, fd); + + assertThat(result).isEqualTo(testCase.expected); + } + + @Test + public void hasSingleField_emptyBytes_returnsFalse() throws Exception { + FieldLiteDescriptor singleInt64Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", TestAllTypes.SINGLE_INT64_FIELD_NUMBER) + .get(); + + boolean result = PROTO_LITE_CEL_VALUE_CONVERTER.hasSingleField(ByteString.EMPTY, singleInt64Fd); + + assertThat(result).isFalse(); + } + + @Test + public void hasSingleField_emptyPackedRepeatedField_returnsFalse() throws Exception { + FieldLiteDescriptor repeatedInt32Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", + TestAllTypes.REPEATED_INT32_FIELD_NUMBER) + .get(); + ByteArrayOutputStream emptyPackedOut = new ByteArrayOutputStream(); + CodedOutputStream emptyPackedCos = CodedOutputStream.newInstance(emptyPackedOut); + emptyPackedCos.writeInt64(TestAllTypes.SINGLE_INT64_FIELD_NUMBER, 42L); + emptyPackedCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + emptyPackedCos.flush(); + ByteString emptyPackedBytes = ByteString.copyFrom(emptyPackedOut.toByteArray()); + + boolean result = + PROTO_LITE_CEL_VALUE_CONVERTER.hasSingleField(emptyPackedBytes, repeatedInt32Fd); + + assertThat(result).isFalse(); + } + + @Test + public void hasSingleField_emptyPackedFollowedByPopulatedPacked_returnsTrue() throws Exception { + FieldLiteDescriptor repeatedInt32Fd = + PROTO_LITE_CEL_VALUE_CONVERTER + .findFieldDescriptor( + "cel.expr.conformance.proto3.TestAllTypes", + TestAllTypes.REPEATED_INT32_FIELD_NUMBER) + .get(); + ByteArrayOutputStream emptyThenPopulatedOut = new ByteArrayOutputStream(); + CodedOutputStream emptyThenPopulatedCos = CodedOutputStream.newInstance(emptyThenPopulatedOut); + emptyThenPopulatedCos.writeByteArray(TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[0]); + emptyThenPopulatedCos.writeByteArray( + TestAllTypes.REPEATED_INT32_FIELD_NUMBER, new byte[] {1, 2}); + emptyThenPopulatedCos.flush(); + ByteString emptyThenPopulatedBytes = ByteString.copyFrom(emptyThenPopulatedOut.toByteArray()); + + boolean result = + PROTO_LITE_CEL_VALUE_CONVERTER.hasSingleField(emptyThenPopulatedBytes, repeatedInt32Fd); + + assertThat(result).isTrue(); + } + + @Test + public void skipWireField_groupWireType_throwsUnsupportedOperationException() { + int startGroupTag = (1 << 3) | WireFormat.WIRETYPE_START_GROUP; + + assertThrows( + UnsupportedOperationException.class, + () -> + ProtoLiteCelValueConverter.skipWireField( + startGroupTag, ByteString.EMPTY.newCodedInput())); + } } diff --git a/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java b/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java index 822de79b3..41c9fe666 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java @@ -225,6 +225,49 @@ public void create_withCorruptByteString_throwsOnSelect() { assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); } + @Test + public void create_withCorruptByteString_throwsOnSelectByFieldNumber() { + ByteString corruptBytes = ByteString.copyFrom(new byte[] {0x10, (byte) 0x80}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + corruptBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + SelectField selectField = SelectField.create(2L, "single_int64", 3, 0L); + + IllegalArgumentException thrown = + assertThrows( + IllegalArgumentException.class, + () -> messageLiteValue.selectByFieldNumber(selectField)); + + assertThat(thrown) + .hasMessageThat() + .contains( + "Failed to decode proto message of type: cel.expr.conformance.proto3.TestAllTypes"); + assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); + } + + @Test + public void create_withCorruptByteString_throwsOnHasFieldByNumber() { + ByteString corruptBytes = ByteString.copyFrom(new byte[] {0x10, (byte) 0x80}); + ProtoMessageLiteValue messageLiteValue = + ProtoMessageLiteValue.create( + corruptBytes, + "cel.expr.conformance.proto3.TestAllTypes", + PROTO_LITE_CEL_VALUE_CONVERTER); + SelectField selectField = SelectField.create(2L, "single_int64"); + + IllegalArgumentException thrown = + assertThrows( + IllegalArgumentException.class, () -> messageLiteValue.hasFieldByNumber(selectField)); + + assertThat(thrown) + .hasMessageThat() + .contains( + "Failed to decode proto message of type: cel.expr.conformance.proto3.TestAllTypes"); + assertThat(thrown).hasCauseThat().isInstanceOf(IOException.class); + } + @SuppressWarnings("ImmutableEnumChecker") // Test only private enum SelectFieldTestCase { BOOL("single_bool", true),