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 fe114cdd0..a154d9603 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java +++ b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java @@ -16,6 +16,7 @@ import static com.google.common.base.Preconditions.checkNotNull; +import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Defaults; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; @@ -29,6 +30,7 @@ import com.google.protobuf.MessageLite; import com.google.protobuf.WireFormat; import dev.cel.common.annotations.Internal; +import dev.cel.common.exceptions.CelAttributeNotFoundException; import dev.cel.common.internal.CelLiteDescriptorPool; import dev.cel.common.internal.WellKnownProto; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; @@ -59,6 +61,8 @@ public final class ProtoLiteCelValueConverter extends BaseProtoCelValueConverter { private static final String MAP_KEY_FIELD_NAME = "key"; private static final String MAP_VALUE_FIELD_NAME = "value"; + private static final int MAP_KEY_FIELD_NUMBER = 1; + private static final int MAP_VALUE_FIELD_NUMBER = 2; private final CelLiteDescriptorPool descriptorPool; @@ -194,21 +198,6 @@ private static MessageLite mergeMessageLite( } } - 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); @@ -284,12 +273,10 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) { } 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); + CodedInputStream inputStream, + FieldLiteDescriptor keyDescriptor, + FieldLiteDescriptor valueDescriptor) + throws IOException { int length = inputStream.readInt32(); int oldLimit = inputStream.pushLimit(length); Object key = null; @@ -337,14 +324,9 @@ private Map.Entry readSingleMapEntry( } boolean hasSingleField(ByteString bytes, FieldLiteDescriptor fieldDescriptor) throws IOException { - return hasSingleField( - bytes, - fieldDescriptor.getFieldNumber(), - fieldDescriptor.getEncodingType().equals(EncodingType.LIST) && isPackable(fieldDescriptor)); - } - - static boolean hasSingleField(ByteString bytes, int targetFieldNumber, boolean isPackableList) - throws IOException { + int targetFieldNumber = fieldDescriptor.getFieldNumber(); + boolean isPackableList = + fieldDescriptor.getEncodingType().equals(EncodingType.LIST) && isPackable(fieldDescriptor); CodedInputStream inputStream = bytes.newCodedInput(); for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { int fieldNumber = WireFormat.getTagFieldNumber(tag); @@ -370,6 +352,126 @@ static boolean hasSingleField(ByteString bytes, int targetFieldNumber, boolean i return false; } + /** + * Selects {@code field} from {@code bytes}, decoding it by the type information in {@code field} + * for fields missing from the descriptor pool. Returns the field's default value if it's absent. + * Throws {@link CelAttributeNotFoundException} if {@code field} has no type code. + */ + Object selectByFieldNumber(ByteString bytes, SelectField field) throws IOException { + if (field.typeCode() == SelectField.NO_TYPE_CODE) { + throw CelAttributeNotFoundException.forFieldResolution(field.fieldName()); + } + Object fieldValue = readFieldByNumber(bytes, field); + if (fieldValue != null) { + return fieldValue; + } + if (field.defaultValue() != null) { + return field.defaultValue(); + } + return getDefaultCelValue(newFieldDescriptor(field)); + } + + /** + * Finds {@code field} in {@code bytes}, decoding it by the type information in {@code field} for + * fields missing from the descriptor pool. A field without a type code is decoded as a message. + */ + Optional findByFieldNumber(ByteString bytes, SelectField field) throws IOException { + return Optional.ofNullable(readFieldByNumber(bytes, field)); + } + + /** Returns whether {@code field} is present in {@code bytes}. */ + boolean hasFieldByNumber(ByteString bytes, SelectField field) throws IOException { + return hasSingleField(bytes, newFieldDescriptor(field)); + } + + private @Nullable Object readFieldByNumber(ByteString bytes, SelectField field) + throws IOException { + FieldLiteDescriptor fieldDescriptor = newFieldDescriptor(field); + SelectField.MapEntrySpec mapEntrySpec = field.mapEntrySpec(); + if (mapEntrySpec == null) { + return readSingleField(bytes, fieldDescriptor); + } + // The map entry type has no descriptor either, so its key and value are described by the spec. + FieldLiteDescriptor keyDescriptor = + newFieldDescriptor( + MAP_KEY_FIELD_NUMBER, + MAP_KEY_FIELD_NAME, + EncodingType.SINGULAR, + mapEntrySpec.keyTypeCode(), + /* protoTypeName= */ ""); + FieldLiteDescriptor valueDescriptor = + newFieldDescriptor( + MAP_VALUE_FIELD_NUMBER, + MAP_VALUE_FIELD_NAME, + EncodingType.SINGULAR, + mapEntrySpec.valueTypeCode(), + field.protoTypeName()); + CodedInputStream inputStream = bytes.newCodedInput(); + Map mapValues = null; + for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { + if (WireFormat.getTagFieldNumber(tag) != field.fieldNumber()) { + skipWireField(tag, inputStream); + continue; + } + mapValues = + readMapField( + WireFormat.getTagWireType(tag), + inputStream, + fieldDescriptor, + keyDescriptor, + valueDescriptor, + mapValues); + } + return mapValues == null ? null : resolveFieldValue(finalizeFieldValue(mapValues)); + } + + /** Describes {@code field} by the type information it carries. */ + private static FieldLiteDescriptor newFieldDescriptor(SelectField field) { + if (field.mapEntrySpec() != null) { + // SelectField doesn't name the map entry type; its protoTypeName() is the map's value type. + return newFieldDescriptor( + field.fieldNumber(), + field.fieldName(), + EncodingType.MAP, + SelectField.MESSAGE_TYPE_CODE, + /* protoTypeName= */ ""); + } + if (field.typeCode() == SelectField.NO_TYPE_CODE) { + // Only presence tests omit the type code. Their presence check doesn't depend on the type, + // and the fields they navigate through are always messages. + return newFieldDescriptor( + field.fieldNumber(), + field.fieldName(), + EncodingType.SINGULAR, + SelectField.MESSAGE_TYPE_CODE, + /* protoTypeName= */ ""); + } + return newFieldDescriptor( + field.fieldNumber(), + field.fieldName(), + field.defaultValue() instanceof List ? EncodingType.LIST : EncodingType.SINGULAR, + field.typeCode(), + field.protoTypeName()); + } + + private static FieldLiteDescriptor newFieldDescriptor( + int fieldNumber, + String fieldName, + EncodingType encodingType, + int typeCode, + String protoTypeName) { + FieldLiteDescriptor.Type protoFieldType = FieldLiteDescriptor.Type.forNumber(typeCode); + return new FieldLiteDescriptor( + fieldNumber, + fieldName, + // FieldLiteDescriptor.JavaType's constants match WireFormat.JavaType's by name. + JavaType.valueOf(protoFieldType.toWireFormatFieldType().getJavaType().name()), + encodingType, + protoFieldType, + /* isPacked= */ false, + protoTypeName); + } + /** * Decodes every known field in {@code bytes}, keyed by field name. Each value must be passed to * {@link #resolveFieldValue} to obtain its CEL value. @@ -520,18 +622,38 @@ private Object readSingularField( return repeatedValues; } + private Map readMapField( + int tagWireType, + CodedInputStream inputStream, + FieldLiteDescriptor fieldDescriptor, + @Nullable Object existingValue) + throws IOException { + MessageLiteDescriptor entryDescriptor = + descriptorPool.getDescriptorOrThrow(fieldDescriptor.getFieldProtoTypeName()); + return readMapField( + tagWireType, + inputStream, + fieldDescriptor, + entryDescriptor.getByFieldNameOrThrow(MAP_KEY_FIELD_NAME), + entryDescriptor.getByFieldNameOrThrow(MAP_VALUE_FIELD_NAME), + existingValue); + } + // Safe because MAP fields only ever store a LinkedHashMap as their accumulated value. @SuppressWarnings("unchecked") private Map readMapField( int tagWireType, CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor, + FieldLiteDescriptor keyDescriptor, + FieldLiteDescriptor valueDescriptor, @Nullable Object existingValue) throws IOException { checkWireType(tagWireType, fieldDescriptor); Map mapValues = existingValue != null ? (Map) existingValue : new LinkedHashMap<>(); - Map.Entry mapEntry = readSingleMapEntry(inputStream, fieldDescriptor); + Map.Entry mapEntry = + readSingleMapEntry(inputStream, keyDescriptor, valueDescriptor); mapValues.put(mapEntry.getKey(), mapEntry.getValue()); return mapValues; } @@ -565,6 +687,7 @@ private static boolean isPackable(FieldLiteDescriptor fieldDescriptor) { return fieldDescriptor.getProtoFieldType().toWireFormatFieldType().isPackable(); } + @VisibleForTesting static void skipWireField(int tag, CodedInputStream inputStream) throws IOException { int tagWireType = WireFormat.getTagWireType(tag); switch (tagWireType) { @@ -582,25 +705,6 @@ static void skipWireField(int tag, CodedInputStream inputStream) throws IOExcept } } - static Object readUnknownField(int tagWireType, CodedInputStream inputStream) throws IOException { - switch (tagWireType) { - case WireFormat.WIRETYPE_VARINT: - return inputStream.readInt64(); - case WireFormat.WIRETYPE_FIXED64: - return inputStream.readFixed64(); - case WireFormat.WIRETYPE_LENGTH_DELIMITED: - return inputStream.readBytes(); - case WireFormat.WIRETYPE_FIXED32: - return inputStream.readFixed32(); - case WireFormat.WIRETYPE_START_GROUP: - case WireFormat.WIRETYPE_END_GROUP: - // TODO: Support groups - throw new UnsupportedOperationException("Groups are not supported"); - default: - throw new IllegalArgumentException("Unknown wire type: " + tagWireType); - } - } - /** * A field value holding well-known type messages, whose conversion to CEL values is deferred * until {@link #resolveFieldValue}. 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 e22e9f552..fb2af1cb0 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java @@ -154,10 +154,12 @@ public Object selectByFieldNumber(SelectField field) { } return protoLiteCelValueConverter().getDefaultCelValue(fd); } - return RawProtoMessageLiteValue.selectWireOrDefault( - field, - RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()), - protoLiteCelValueConverter()); + try { + return protoLiteCelValueConverter().selectByFieldNumber(toByteString(), field); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } } @Override @@ -166,7 +168,12 @@ public boolean hasFieldByNumber(SelectField field) { if (fd != null) { return hasField(fd); } - return RawProtoMessageLiteValue.isPresentInWire(toByteString(), field); + try { + return protoLiteCelValueConverter().hasFieldByNumber(toByteString(), field); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } } @Override @@ -175,10 +182,12 @@ public Optional findByFieldNumber(SelectField field) { if (fd != null) { return Optional.ofNullable(readField(fd)); } - return RawProtoMessageLiteValue.navigateWire( - field, - RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()), - protoLiteCelValueConverter()); + try { + return protoLiteCelValueConverter().findByFieldNumber(toByteString(), field); + } catch (IOException e) { + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); + } } private @Nullable Object readField(FieldLiteDescriptor fd) { 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 848dd3008..77cae2c43 100644 --- a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java @@ -17,26 +17,15 @@ import static com.google.common.base.Preconditions.checkNotNull; import com.google.auto.value.AutoValue; -import com.google.common.collect.ImmutableCollection; -import com.google.common.collect.ImmutableList; -import com.google.common.collect.ImmutableMap; -import com.google.common.collect.Iterables; -import com.google.common.primitives.UnsignedLong; import com.google.errorprone.annotations.Immutable; import com.google.protobuf.ByteString; import com.google.protobuf.CodedInputStream; -import com.google.protobuf.WireFormat; import dev.cel.common.exceptions.CelAttributeNotFoundException; import dev.cel.common.types.CelType; import dev.cel.common.types.StructTypeReference; -import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; import java.io.IOException; -import java.util.AbstractMap; -import java.util.List; import java.util.Locale; -import java.util.Map; import java.util.Optional; -import org.jspecify.annotations.Nullable; /** * RawProtoMessageLiteValue enables descriptorless evaluation of protobuf messages to address @@ -53,8 +42,6 @@ abstract class RawProtoMessageLiteValue extends StructValue findByFieldNumber(SelectField field) { - return navigateWire( - 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)); - } + return protoLiteCelValueConverter().selectByFieldNumber(toByteString(), field); } catch (IOException e) { - throw new IllegalStateException("Failed to parse raw proto message wire bytes", e); - } - return entries.build(); - } - - /** - * Decodes a field value from preserved wire bytes, falling back to default values. - * - *

Package-private: shared with {@code ProtoMessageLiteValue} for unknown field resolution. - */ - static Object selectWireOrDefault( - SelectField field, ImmutableList unknowns, ProtoLiteCelValueConverter converter) { - if (unknowns.isEmpty()) { - return resolveDefault(field, converter); - } - return decodeWireField(field, unknowns, converter); - } - - private static Object decodeWireField( - SelectField field, ImmutableList unknowns, ProtoLiteCelValueConverter converter) { - SelectField.MapEntrySpec mapEntrySpec = field.mapEntrySpec(); - if (mapEntrySpec != null) { - return decodeMapEntries( - unknowns, mapEntrySpec, field.protoTypeName(), field.fieldName(), converter); - } - - int typeCode = field.typeCode(); - if (typeCode == SelectField.NO_TYPE_CODE) { - throw CelAttributeNotFoundException.forFieldResolution(field.fieldName()); - } - - boolean isRepeated = field.defaultValue() instanceof List; - - return decodeWireEntries(unknowns, typeCode, field.protoTypeName(), isRepeated, converter); - } - - private static Object resolveDefault(SelectField field, ProtoLiteCelValueConverter converter) { - if (field.defaultValue() != null) { - return field.defaultValue(); - } - return decodeMessageValue(ByteString.EMPTY, field.protoTypeName(), converter); - } - - /** - * 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(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; - } - - // 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 (!isPackableRepeated(field)) { - return true; - } - - for (Object raw : unknowns) { - if (!(raw instanceof ByteString) || !((ByteString) raw).isEmpty()) { - return true; - } - } - 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. - * - *

Package-private: shared with {@code ProtoMessageLiteValue} for unknown field resolution. - */ - static Optional navigateWire( - SelectField field, ImmutableList unknowns, ProtoLiteCelValueConverter converter) { - if (!isPresentInWire(field, unknowns)) { - return Optional.empty(); - } - if (field.typeCode() != SelectField.NO_TYPE_CODE) { - return Optional.of(selectWireOrDefault(field, unknowns, converter)); - } - Object lastEntry = unknowns.get(unknowns.size() - 1); - if (lastEntry instanceof ByteString) { - return Optional.of( - decodeWireEntries( - unknowns, - FieldLiteDescriptor.Type.MESSAGE.getNumber(), - UNKNOWN_MESSAGE_TYPE_NAME, - /* isRepeated= */ false, - converter)); + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); } - return Optional.of(lastEntry); } - /** - * Decodes map entries from wire bytes without descriptors (version skew), using the - * optimizer-provided {@link SelectField.MapEntrySpec}. Per protobuf semantics, repeated key or - * value tags within an entry are last-one-wins for scalars and merged for messages, and absent - * keys or values take their type's default. - */ - private static ImmutableMap decodeMapEntries( - ImmutableList unknowns, - SelectField.MapEntrySpec spec, - String valueProtoTypeName, - String fieldName, - ProtoLiteCelValueConverter converter) { - ImmutableMap.Builder mapBuilder = ImmutableMap.builder(); + @Override + public boolean hasFieldByNumber(SelectField field) { try { - for (Object raw : unknowns) { - ByteString bytes = requireType(raw, ByteString.class, WireFormat.FieldType.MESSAGE); - mapBuilder.put( - decodeSingleMapEntry(bytes.newCodedInput(), spec, valueProtoTypeName, converter)); - } + return protoLiteCelValueConverter().hasFieldByNumber(toByteString(), field); } catch (IOException e) { - throw new IllegalArgumentException("Failed to decode map entry for field: " + fieldName, e); - } - return mapBuilder.buildKeepingLast(); - } - - private static Map.Entry decodeSingleMapEntry( - CodedInputStream in, - SelectField.MapEntrySpec spec, - String valueProtoTypeName, - ProtoLiteCelValueConverter converter) - throws IOException { - WireFormat.FieldType keyWireType = - FieldLiteDescriptor.Type.forNumber(spec.keyTypeCode()).toWireFormatFieldType(); - WireFormat.FieldType valWireType = - FieldLiteDescriptor.Type.forNumber(spec.valueTypeCode()).toWireFormatFieldType(); - boolean isMessageValue = valWireType == WireFormat.FieldType.MESSAGE; - Object key = resolveDefaultMapScalarValue(spec.keyTypeCode()); - Object value = - isMessageValue ? ByteString.EMPTY : resolveDefaultMapScalarValue(spec.valueTypeCode()); - for (int tag = in.readTag(); tag != 0; tag = in.readTag()) { - int tagWireType = WireFormat.getTagWireType(tag); - int fieldNumber = WireFormat.getTagFieldNumber(tag); - Object parsedValue = ProtoLiteCelValueConverter.readUnknownField(tagWireType, in); - switch (fieldNumber) { - case MAP_KEY_FIELD_NUMBER: - key = decodeWireValue(parsedValue, keyWireType, /* protoTypeName= */ "", converter); - break; - case MAP_VALUE_FIELD_NUMBER: - value = - isMessageValue - ? ((ByteString) value) - .concat(requireType(parsedValue, ByteString.class, valWireType)) - : decodeWireValue(parsedValue, valWireType, /* protoTypeName= */ "", converter); - break; - default: - throw new IllegalStateException("Unexpected field number in map entry: " + fieldNumber); - } - } - if (isMessageValue) { - value = decodeWireValue(value, valWireType, valueProtoTypeName, converter); - } - return new AbstractMap.SimpleImmutableEntry<>(key, value); - } - - private static Object resolveDefaultMapScalarValue(int typeCode) { - FieldLiteDescriptor.Type protoType = FieldLiteDescriptor.Type.forNumber(typeCode); - switch (protoType) { - case BOOL: - return false; - case INT32: - case INT64: - case SINT32: - case SINT64: - case SFIXED32: - case SFIXED64: - case ENUM: - return 0L; - case UINT32: - case UINT64: - case FIXED32: - case FIXED64: - return UnsignedLong.ZERO; - case FLOAT: - case DOUBLE: - return 0.0d; - case STRING: - return ""; - case BYTES: - return CelByteString.EMPTY; - default: - throw new IllegalArgumentException("Unsupported map scalar type code: " + typeCode); - } - } - - static @Nullable Object decodeWireEntries( - ImmutableCollection entries, - int typeCode, - String protoTypeName, - boolean isRepeated, - ProtoLiteCelValueConverter converter) { - WireFormat.FieldType fieldType = - FieldLiteDescriptor.Type.forNumber(typeCode).toWireFormatFieldType(); - if (fieldType == WireFormat.FieldType.GROUP) { - throw new UnsupportedOperationException("Groups are not supported"); - } - if (entries.isEmpty()) { - return isRepeated ? ImmutableList.of() : null; - } - if (isRepeated) { - ImmutableList.Builder listBuilder = ImmutableList.builder(); - for (Object raw : entries) { - if (fieldType.isPackable() && (raw instanceof ByteString)) { - listBuilder.addAll(decodePacked((ByteString) raw, fieldType)); - } else { - listBuilder.add(decodeWireValue(raw, fieldType, protoTypeName, converter)); - } - } - return listBuilder.build(); - } - if (fieldType == WireFormat.FieldType.MESSAGE) { - ByteString mergedBytes = ByteString.EMPTY; - for (Object item : entries) { - mergedBytes = mergedBytes.concat(requireType(item, ByteString.class, fieldType)); - } - return decodeWireValue(mergedBytes, fieldType, protoTypeName, converter); - } - // Protobuf "last one wins" semantics for non-repeated scalar fields - return decodeWireValue(Iterables.getLast(entries), fieldType, protoTypeName, converter); - } - - static Object decodeWireValue( - Object raw, - WireFormat.FieldType fieldType, - String protoTypeName, - ProtoLiteCelValueConverter converter) { - switch (fieldType) { - case DOUBLE: - return Double.longBitsToDouble(requireType(raw, Long.class, fieldType)); - case FLOAT: - return (double) Float.intBitsToFloat(requireType(raw, Integer.class, fieldType)); - case INT64: - case SFIXED64: - return requireType(raw, Long.class, fieldType); - case INT32: - case ENUM: - return (long) requireType(raw, Long.class, fieldType).intValue(); - case UINT64: - case FIXED64: - return UnsignedLong.fromLongBits(requireType(raw, Long.class, fieldType)); - case FIXED32: - return UnsignedLong.fromLongBits( - Integer.toUnsignedLong(requireType(raw, Integer.class, fieldType))); - case BOOL: - return requireType(raw, Long.class, fieldType) != 0L; - case STRING: - ByteString stringBytes = requireType(raw, ByteString.class, fieldType); - if (!stringBytes.isValidUtf8()) { - throw new IllegalArgumentException("Invalid UTF-8 in string field"); - } - return stringBytes.toStringUtf8(); - case GROUP: - throw new UnsupportedOperationException("Groups are not supported"); - case MESSAGE: - ByteString msgBytes = requireType(raw, ByteString.class, fieldType); - return decodeMessageValue(msgBytes, protoTypeName, converter); - case BYTES: - return CelByteString.of(requireType(raw, ByteString.class, fieldType).toByteArray()); - case UINT32: - return UnsignedLong.fromLongBits(requireType(raw, Long.class, fieldType) & 0xFFFFFFFFL); - case SFIXED32: - return (long) requireType(raw, Integer.class, fieldType); - case SINT32: - return (long) - CodedInputStream.decodeZigZag32(requireType(raw, Long.class, fieldType).intValue()); - case SINT64: - return CodedInputStream.decodeZigZag64(requireType(raw, Long.class, fieldType)); - } - 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)) { throw new IllegalArgumentException( - String.format( - "Expected %s for wire type %s, but got: %s", - expectedType.getSimpleName(), - fieldType, - raw != null ? raw.getClass().getName() : "null")); + "Failed to decode proto message of type: " + celType().name(), e); } - return expectedType.cast(raw); } - private static ImmutableList decodePacked( - ByteString bytes, WireFormat.FieldType fieldType) { + @Override + public Optional findByFieldNumber(SelectField field) { try { - CodedInputStream in = bytes.newCodedInput(); - ImmutableList.Builder builder = ImmutableList.builder(); - while (!in.isAtEnd()) { - switch (fieldType) { - case DOUBLE: - builder.add(Double.longBitsToDouble(in.readFixed64())); - break; - case FLOAT: - builder.add((double) Float.intBitsToFloat(in.readFixed32())); - break; - case INT64: - builder.add(in.readInt64()); - break; - case UINT64: - builder.add(UnsignedLong.fromLongBits(in.readUInt64())); - break; - case INT32: - builder.add((long) in.readInt32()); - break; - case FIXED64: - builder.add(UnsignedLong.fromLongBits(in.readFixed64())); - break; - case FIXED32: - builder.add(UnsignedLong.fromLongBits(Integer.toUnsignedLong(in.readFixed32()))); - break; - case BOOL: - builder.add(in.readBool()); - break; - case UINT32: - builder.add(UnsignedLong.fromLongBits(Integer.toUnsignedLong(in.readUInt32()))); - break; - case ENUM: - builder.add((long) in.readEnum()); - break; - case SFIXED32: - builder.add((long) in.readSFixed32()); - break; - case SFIXED64: - builder.add(in.readSFixed64()); - break; - case SINT32: - builder.add((long) in.readSInt32()); - break; - case SINT64: - builder.add(in.readSInt64()); - break; - default: - throw new IllegalArgumentException("Unsupported packed proto field type: " + fieldType); - } - } - return builder.build(); + return protoLiteCelValueConverter().findByFieldNumber(toByteString(), field); } catch (IOException e) { - throw new IllegalStateException("Failed to parse packed repeated field", e); + throw new IllegalArgumentException( + "Failed to decode proto message of type: " + celType().name(), e); } } 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 efa7cb082..69c4eba81 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java @@ -50,7 +50,6 @@ 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.Map; import java.util.NoSuchElementException; @@ -369,104 +368,6 @@ public void readAllFields_nestedMessageWithoutDescriptor_returnsRawProtoMessageL .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } - @Test - public void tryDecodeProtoMessage_wellKnownType_returnsDecodedValue() { - Int32Value int32Value = Int32Value.of(42); - - Optional decoded = - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( - int32Value.toByteString(), "google.protobuf.Int32Value"); - - assertThat(decoded).hasValue(42L); - } - - @Test - 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( - nestedMsg, - "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", - 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() { - 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.tryDecodeProtoMessage(ByteString.EMPTY, "google.protobuf.Int32Value"); - - assertThat(decoded).isEmpty(); - } - - @Test - 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.tryDecodeProtoMessage( - corruptBytes, "google.protobuf.Int32Value")); - - assertThat(exception) - .hasMessageThat() - .contains("Failed to decode proto message of type: google.protobuf.Int32Value"); - assertThat(exception).hasCauseThat().isInstanceOf(IOException.class); - } - - @Test - public void tryDecodeProtoMessage_anyType_throwsUnsupportedOperationException() { - UnsupportedOperationException exception = - assertThrows( - UnsupportedOperationException.class, - () -> - 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(); 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 fced65916..b9cc6382a 100644 --- a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java @@ -15,10 +15,8 @@ package dev.cel.common.values; import static com.google.common.truth.Truth.assertThat; -import static java.nio.charset.StandardCharsets.UTF_8; import static org.junit.Assert.assertThrows; -import com.google.common.collect.ImmutableCollection; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; @@ -51,17 +49,6 @@ public final class RawProtoMessageLiteValueTest { private static final ProtoLiteCelValueConverter EMPTY_CONVERTER = ProtoLiteCelValueConverter.newInstance(DefaultLiteDescriptorPool.newInstance()); - private static Object decodeWireEntries( - ImmutableCollection entries, int typeCode, String protoTypeName, boolean isRepeated) { - return RawProtoMessageLiteValue.decodeWireEntries( - entries, typeCode, protoTypeName, isRepeated, EMPTY_CONVERTER); - } - - private static Object decodeWireValue( - Object raw, WireFormat.FieldType fieldType, String protoTypeName) { - return RawProtoMessageLiteValue.decodeWireValue(raw, fieldType, protoTypeName, EMPTY_CONVERTER); - } - @Test public void create_accessorsAndType() { ByteString bytes = ByteString.copyFromUtf8("test"); @@ -222,14 +209,10 @@ public void breakingFieldChange_varintWireReadAsString_hasFieldReturnsTrueButSel IllegalArgumentException thrownSelect = assertThrows( IllegalArgumentException.class, () -> value.selectByFieldNumber(breakingField)); - assertThat(thrownSelect) - .hasMessageThat() - .contains("Expected ByteString for wire type STRING, but got: java.lang.Long"); + assertThat(thrownSelect).hasCauseThat().hasMessageThat().contains("unexpected wire type"); IllegalArgumentException thrownFind = assertThrows(IllegalArgumentException.class, () -> value.findByFieldNumber(breakingField)); - assertThat(thrownFind) - .hasMessageThat() - .contains("Expected ByteString for wire type STRING, but got: java.lang.Long"); + assertThat(thrownFind).hasCauseThat().hasMessageThat().contains("unexpected wire type"); } @Test @@ -255,598 +238,10 @@ public void breakingFieldChange_varintWireReadAsString_hasFieldReturnsTrueButSel IllegalArgumentException thrownSelect = assertThrows( IllegalArgumentException.class, () -> value.selectByFieldNumber(breakingField)); - assertThat(thrownSelect).hasMessageThat().contains("Expected Long for wire type INT64"); + assertThat(thrownSelect).hasCauseThat().hasMessageThat().contains("unexpected wire type"); IllegalArgumentException thrownFind = assertThrows(IllegalArgumentException.class, () -> value.findByFieldNumber(breakingField)); - assertThat(thrownFind).hasMessageThat().contains("Expected Long for wire type INT64"); - } - - @Test - public void decodeWireEntries_emptySingularEntries_returnsNull() { - Object intResult = - decodeWireEntries( - ImmutableList.of(), - FieldLiteDescriptor.Type.INT64.getNumber(), - "custom.Message", - /* isRepeated= */ false); - Object messageResult = - decodeWireEntries( - ImmutableList.of(), - FieldLiteDescriptor.Type.MESSAGE.getNumber(), - "custom.Message", - /* isRepeated= */ false); - - assertThat(intResult).isNull(); - assertThat(messageResult).isNull(); - } - - @Test - public void decodeWireEntries_emptyRepeatedEntries_returnsEmptyList() { - Object result = - decodeWireEntries( - ImmutableList.of(), - FieldLiteDescriptor.Type.INT64.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat((Iterable) result).isEmpty(); - } - - @Test - public void decodeWireEntries_nonRepeated_lastOneWins() { - Object decoded = - decodeWireEntries( - ImmutableList.of(10L, 20L, 30L), - FieldLiteDescriptor.Type.INT64.getNumber(), - "custom.Message", - /* isRepeated= */ false); - - assertThat(decoded).isEqualTo(30L); - } - - @Test - public void decodeWireEntries_repeatedUnpacked() { - Object decoded = - decodeWireEntries( - ImmutableList.of(10L, 20L, 30L), - FieldLiteDescriptor.Type.INT64.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded).isEqualTo(ImmutableList.of(10L, 20L, 30L)); - } - - @Test - public void decodeWireEntries_packedInt32() throws Exception { - ByteArrayOutputStream baos = new ByteArrayOutputStream(); - CodedOutputStream cos = CodedOutputStream.newInstance(baos); - cos.writeInt32NoTag(1); - cos.writeInt32NoTag(2); - cos.writeInt32NoTag(3); - cos.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray())), - FieldLiteDescriptor.Type.INT32.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded).isEqualTo(ImmutableList.of(1L, 2L, 3L)); - } - - @Test - public void decodeWireEntries_packedInt64() throws Exception { - ByteArrayOutputStream baos = new ByteArrayOutputStream(); - CodedOutputStream cos = CodedOutputStream.newInstance(baos); - cos.writeInt64NoTag(100L); - cos.writeInt64NoTag(200L); - cos.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray())), - FieldLiteDescriptor.Type.INT64.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded).isEqualTo(ImmutableList.of(100L, 200L)); - } - - @Test - public void decodeWireEntries_packedUint32() throws Exception { - ByteArrayOutputStream baos = new ByteArrayOutputStream(); - CodedOutputStream cos = CodedOutputStream.newInstance(baos); - cos.writeUInt32NoTag(50); - cos.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray())), - FieldLiteDescriptor.Type.UINT32.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded).isEqualTo(ImmutableList.of(UnsignedLong.fromLongBits(50L))); - } - - @Test - public void decodeWireEntries_packedUint64() throws Exception { - ByteArrayOutputStream baos = new ByteArrayOutputStream(); - CodedOutputStream cos = CodedOutputStream.newInstance(baos); - cos.writeUInt64NoTag(999L); - cos.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray())), - FieldLiteDescriptor.Type.UINT64.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded).isEqualTo(ImmutableList.of(UnsignedLong.fromLongBits(999L))); - } - - @Test - public void decodeWireEntries_packedSint32AndSint64() throws Exception { - ByteArrayOutputStream baos32 = new ByteArrayOutputStream(); - CodedOutputStream cos32 = CodedOutputStream.newInstance(baos32); - cos32.writeSInt32NoTag(-10); - cos32.writeSInt32NoTag(20); - cos32.flush(); - - Object decoded32 = - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos32.toByteArray())), - FieldLiteDescriptor.Type.SINT32.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded32).isEqualTo(ImmutableList.of(-10L, 20L)); - - ByteArrayOutputStream baos64 = new ByteArrayOutputStream(); - CodedOutputStream cos64 = CodedOutputStream.newInstance(baos64); - cos64.writeSInt64NoTag(-100L); - cos64.writeSInt64NoTag(200L); - cos64.flush(); - - Object decoded64 = - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos64.toByteArray())), - FieldLiteDescriptor.Type.SINT64.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat(decoded64).isEqualTo(ImmutableList.of(-100L, 200L)); - } - - @Test - public void decodeWireEntries_packedFixedAndSFixed() throws Exception { - ByteArrayOutputStream baos = new ByteArrayOutputStream(); - CodedOutputStream cos = CodedOutputStream.newInstance(baos); - cos.writeFixed32NoTag(10); - cos.writeFixed64NoTag(20L); - cos.writeSFixed32NoTag(-30); - cos.writeSFixed64NoTag(-40L); - cos.flush(); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray()).substring(0, 4)), - FieldLiteDescriptor.Type.FIXED32.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(UnsignedLong.fromLongBits(10L))); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray()).substring(4, 12)), - FieldLiteDescriptor.Type.FIXED64.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(UnsignedLong.fromLongBits(20L))); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray()).substring(12, 16)), - FieldLiteDescriptor.Type.SFIXED32.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(-30L)); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baos.toByteArray()).substring(16, 24)), - FieldLiteDescriptor.Type.SFIXED64.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(-40L)); - } - - @Test - public void decodeWireEntries_packedBoolFloatDoubleEnum() throws Exception { - ByteArrayOutputStream baosBool = new ByteArrayOutputStream(); - CodedOutputStream cosBool = CodedOutputStream.newInstance(baosBool); - cosBool.writeBoolNoTag(true); - cosBool.writeBoolNoTag(false); - cosBool.flush(); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baosBool.toByteArray())), - FieldLiteDescriptor.Type.BOOL.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(true, false)); - - ByteArrayOutputStream baosFloat = new ByteArrayOutputStream(); - CodedOutputStream cosFloat = CodedOutputStream.newInstance(baosFloat); - cosFloat.writeFloatNoTag(1.5f); - cosFloat.flush(); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baosFloat.toByteArray())), - FieldLiteDescriptor.Type.FLOAT.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(1.5d)); - - ByteArrayOutputStream baosDouble = new ByteArrayOutputStream(); - CodedOutputStream cosDouble = CodedOutputStream.newInstance(baosDouble); - cosDouble.writeDoubleNoTag(3.14d); - cosDouble.flush(); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baosDouble.toByteArray())), - FieldLiteDescriptor.Type.DOUBLE.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(3.14d)); - - ByteArrayOutputStream baosEnum = new ByteArrayOutputStream(); - CodedOutputStream cosEnum = CodedOutputStream.newInstance(baosEnum); - cosEnum.writeEnumNoTag(2); - cosEnum.flush(); - - assertThat( - decodeWireEntries( - ImmutableList.of(ByteString.copyFrom(baosEnum.toByteArray())), - FieldLiteDescriptor.Type.ENUM.getNumber(), - "custom.Message", - /* isRepeated= */ true)) - .isEqualTo(ImmutableList.of(2L)); - } - - @Test - public void decodeWireValue_allScalarWireTypes() { - assertThat( - decodeWireValue( - Double.doubleToRawLongBits(2.5d), WireFormat.FieldType.DOUBLE, "custom.Message")) - .isEqualTo(2.5d); - - assertThat( - decodeWireValue( - Float.floatToRawIntBits(1.5f), WireFormat.FieldType.FLOAT, "custom.Message")) - .isEqualTo(1.5d); - - assertThat(decodeWireValue(42L, WireFormat.FieldType.INT64, "custom.Message")).isEqualTo(42L); - - assertThat(decodeWireValue(42L, WireFormat.FieldType.INT32, "custom.Message")).isEqualTo(42L); - - assertThat(decodeWireValue(42L, WireFormat.FieldType.UINT64, "custom.Message")) - .isEqualTo(UnsignedLong.fromLongBits(42L)); - - assertThat(decodeWireValue(42L, WireFormat.FieldType.UINT32, "custom.Message")) - .isEqualTo(UnsignedLong.fromLongBits(42L)); - - assertThat(decodeWireValue(100, WireFormat.FieldType.FIXED32, "custom.Message")) - .isEqualTo(UnsignedLong.fromLongBits(100L)); - - assertThat(decodeWireValue(100L, WireFormat.FieldType.FIXED64, "custom.Message")) - .isEqualTo(UnsignedLong.fromLongBits(100L)); - - assertThat(decodeWireValue(-50, WireFormat.FieldType.SFIXED32, "custom.Message")) - .isEqualTo(-50L); - - assertThat(decodeWireValue(-50L, WireFormat.FieldType.SFIXED64, "custom.Message")) - .isEqualTo(-50L); - - assertThat(decodeWireValue(1L, WireFormat.FieldType.BOOL, "custom.Message")).isEqualTo(true); - - assertThat(decodeWireValue(0L, WireFormat.FieldType.BOOL, "custom.Message")).isEqualTo(false); - - assertThat( - decodeWireValue( - ByteString.copyFromUtf8("hello"), WireFormat.FieldType.STRING, "custom.Message")) - .isEqualTo("hello"); - - assertThat( - decodeWireValue( - ByteString.copyFromUtf8("bytes"), WireFormat.FieldType.BYTES, "custom.Message")) - .isEqualTo(CelByteString.of("bytes".getBytes(UTF_8))); - - assertThat( - decodeWireValue( - 1L, // zigzag 1 -> -1 - WireFormat.FieldType.SINT32, - "custom.Message")) - .isEqualTo(-1L); - - assertThat( - decodeWireValue( - 1L, // zigzag 1 -> -1 - WireFormat.FieldType.SINT64, - "custom.Message")) - .isEqualTo(-1L); - - assertThat(decodeWireValue(3L, WireFormat.FieldType.ENUM, "custom.Message")).isEqualTo(3L); - } - - @Test - public void decodeWireValue_messageType_returnsRawProtoMessageLiteValue() { - Object submessage = - decodeWireValue( - ByteString.copyFromUtf8("raw"), WireFormat.FieldType.MESSAGE, "sub.Message"); - - assertThat(submessage).isInstanceOf(RawProtoMessageLiteValue.class); - assertThat(((RawProtoMessageLiteValue) submessage).celType().name()).isEqualTo("sub.Message"); - } - - @Test - public void decodeWireValue_groupType_throwsUnsupportedOperationException() { - ByteString rawBytes = ByteString.copyFromUtf8("raw"); - - UnsupportedOperationException thrown = - assertThrows( - UnsupportedOperationException.class, - () -> decodeWireValue(rawBytes, WireFormat.FieldType.GROUP, "group.Message")); - - assertThat(thrown).hasMessageThat().contains("Groups are not supported"); - } - - @Test - public void decodeWireEntries_groupType_throwsUnsupportedOperationException() { - ImmutableList rawEntries = ImmutableList.of(); - int groupTypeCode = FieldLiteDescriptor.Type.GROUP.getNumber(); - - UnsupportedOperationException thrown = - assertThrows( - UnsupportedOperationException.class, - () -> - decodeWireEntries( - rawEntries, groupTypeCode, "group.Message", /* isRepeated= */ false)); - - assertThat(thrown).hasMessageThat().contains("Groups are not supported"); - } - - @Test - public void decodeWireEntries_invalidTypeCode_throwsIllegalArgumentException() { - ImmutableList rawEntries = ImmutableList.of(); - - assertThrows( - IllegalArgumentException.class, - () -> decodeWireEntries(rawEntries, 999, "custom.Message", /* isRepeated= */ false)); - } - - @Test - public void decodeWireValue_int32HighBits_truncatedToSigned32Bit() { - Object decodedHigh = - decodeWireValue(0x100000005L, WireFormat.FieldType.INT32, "custom.Message"); - Object decodedNegative = - decodeWireValue(0xFFFFFFFF80000000L, WireFormat.FieldType.INT32, "custom.Message"); - - assertThat(decodedHigh).isEqualTo(5L); - assertThat(decodedNegative).isEqualTo(-2147483648L); - } - - @Test - public void decodeWireValue_enumHighBits_truncatedToSigned32Bit() { - Object decodedHigh = decodeWireValue(0x100000005L, WireFormat.FieldType.ENUM, "custom.Message"); - - assertThat(decodedHigh).isEqualTo(5L); - } - - @Test - public void decodeWireValue_typeMismatch_throwsIllegalArgumentException() { - IllegalArgumentException thrownInt64 = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue("not a long", WireFormat.FieldType.INT64, "custom.Message")); - assertThat(thrownInt64).hasMessageThat().contains("Expected Long for wire type INT64"); - - IllegalArgumentException thrownString = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue(100L, WireFormat.FieldType.STRING, "custom.Message")); - assertThat(thrownString).hasMessageThat().contains("Expected ByteString for wire type STRING"); - - IllegalArgumentException thrownBytes = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue(100L, WireFormat.FieldType.BYTES, "custom.Message")); - assertThat(thrownBytes).hasMessageThat().contains("Expected ByteString for wire type BYTES"); - - IllegalArgumentException thrownMessage = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue(100L, WireFormat.FieldType.MESSAGE, "custom.Message")); - assertThat(thrownMessage) - .hasMessageThat() - .contains("Expected ByteString for wire type MESSAGE"); - - IllegalArgumentException thrownFloat = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue(100L, WireFormat.FieldType.FLOAT, "custom.Message")); - assertThat(thrownFloat).hasMessageThat().contains("Expected Integer for wire type FLOAT"); - - IllegalArgumentException thrownDouble = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue(100, WireFormat.FieldType.DOUBLE, "custom.Message")); - assertThat(thrownDouble).hasMessageThat().contains("Expected Long for wire type DOUBLE"); - } - - @Test - public void decodeWireValue_invalidUtf8String_throwsIllegalArgumentException() { - ByteString invalidUtf8 = ByteString.copyFrom(new byte[] {(byte) 0xC0, (byte) 0xAF}); - - IllegalArgumentException thrown = - assertThrows( - IllegalArgumentException.class, - () -> decodeWireValue(invalidUtf8, WireFormat.FieldType.STRING, "custom.Message")); - assertThat(thrown).hasMessageThat().contains("Invalid UTF-8 in string field"); - } - - @Test - public void decodeWireEntries_multiChunkPackedRepeated() throws Exception { - ByteArrayOutputStream baos1 = new ByteArrayOutputStream(); - CodedOutputStream cos1 = CodedOutputStream.newInstance(baos1); - cos1.writeInt32NoTag(1); - cos1.writeInt32NoTag(2); - cos1.flush(); - - ByteArrayOutputStream baos2 = new ByteArrayOutputStream(); - CodedOutputStream cos2 = CodedOutputStream.newInstance(baos2); - cos2.writeInt32NoTag(3); - cos2.writeInt32NoTag(4); - cos2.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of( - ByteString.copyFrom(baos1.toByteArray()), ByteString.copyFrom(baos2.toByteArray())), - FieldLiteDescriptor.Type.INT32.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat((Iterable) decoded).containsExactly(1L, 2L, 3L, 4L).inOrder(); - } - - @Test - public void decodeWireEntries_mixedPackedAndUnpackedRepeated() throws Exception { - ByteArrayOutputStream baos = new ByteArrayOutputStream(); - CodedOutputStream cos = CodedOutputStream.newInstance(baos); - cos.writeInt32NoTag(2); - cos.writeInt32NoTag(3); - cos.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of(1L, ByteString.copyFrom(baos.toByteArray()), 4L), - FieldLiteDescriptor.Type.INT32.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat((Iterable) decoded).containsExactly(1L, 2L, 3L, 4L).inOrder(); - } - - @Test - public void decodeWireEntries_singularMessage_mergesChunks() throws Exception { - ByteArrayOutputStream baos1 = new ByteArrayOutputStream(); - CodedOutputStream cos1 = CodedOutputStream.newInstance(baos1); - cos1.writeInt64(1, 100L); - cos1.flush(); - - ByteArrayOutputStream baos2 = new ByteArrayOutputStream(); - CodedOutputStream cos2 = CodedOutputStream.newInstance(baos2); - cos2.writeInt64(2, 200L); - cos2.flush(); - - Object decoded = - decodeWireEntries( - ImmutableList.of( - ByteString.copyFrom(baos1.toByteArray()), ByteString.copyFrom(baos2.toByteArray())), - FieldLiteDescriptor.Type.MESSAGE.getNumber(), - "sub.Message", - /* isRepeated= */ false); - - assertThat(decoded).isInstanceOf(RawProtoMessageLiteValue.class); - RawProtoMessageLiteValue rawMessage = (RawProtoMessageLiteValue) decoded; - assertThat(rawMessage.toByteString()) - .isEqualTo( - ByteString.copyFrom(baos1.toByteArray()) - .concat(ByteString.copyFrom(baos2.toByteArray()))); - } - - @Test - public void decodeWireValue_uint32HighBit_correctUnsignedLong() { - Object decoded = decodeWireValue(0xFFFFFFFFL, WireFormat.FieldType.UINT32, "custom.Message"); - - assertThat(decoded).isEqualTo(UnsignedLong.valueOf(4294967295L)); - } - - @Test - public void decodeWireValue_fixed32HighBit_correctUnsignedLong() { - Object decoded = decodeWireValue(-1, WireFormat.FieldType.FIXED32, "custom.Message"); - - assertThat(decoded).isEqualTo(UnsignedLong.valueOf(4294967295L)); - } - - @Test - public void decodeWireEntries_repeatedString() { - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFromUtf8("foo"), ByteString.copyFromUtf8("bar")), - FieldLiteDescriptor.Type.STRING.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat((Iterable) decoded).containsExactly("foo", "bar").inOrder(); - } - - @Test - public void decodeWireEntries_repeatedBytes() { - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFromUtf8("foo"), ByteString.copyFromUtf8("bar")), - FieldLiteDescriptor.Type.BYTES.getNumber(), - "custom.Message", - /* isRepeated= */ true); - - assertThat((Iterable) decoded) - .containsExactly( - CelByteString.of("foo".getBytes(UTF_8)), CelByteString.of("bar".getBytes(UTF_8))) - .inOrder(); - } - - @Test - public void decodeWireEntries_repeatedMessage() { - Object decoded = - decodeWireEntries( - ImmutableList.of(ByteString.copyFromUtf8("msg1"), ByteString.copyFromUtf8("msg2")), - FieldLiteDescriptor.Type.MESSAGE.getNumber(), - "sub.Message", - /* isRepeated= */ true); - - 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 - public void decodeWireEntries_packedTruncated_throwsIllegalStateException() { - // Varint with MSB set (0x80) indicates continuation, but stream ends prematurely. - ByteString truncated = ByteString.copyFrom(new byte[] {(byte) 0x80}); - - IllegalStateException thrown = - assertThrows( - IllegalStateException.class, - () -> - decodeWireEntries( - ImmutableList.of(truncated), - FieldLiteDescriptor.Type.INT32.getNumber(), - "custom.Message", - /* isRepeated= */ true)); - - assertThat(thrown).hasMessageThat().contains("Failed to parse packed repeated field"); + assertThat(thrownFind).hasCauseThat().hasMessageThat().contains("unexpected wire type"); } @Test @@ -1096,17 +491,17 @@ public void selectByFieldNumber_decodesExpectedValue( } @Test - public void findByFieldNumber_scalarFieldWithoutDescriptor_returnsScalar() { + public void findByFieldNumber_scalarFieldWithoutTypeCode_throwsIllegalArgumentException() { TestAllTypes proto = TestAllTypes.newBuilder().setSingleInt64(99L).build(); RawProtoMessageLiteValue raw = RawProtoMessageLiteValue.create( proto.toByteString(), "cel.expr.conformance.proto3.TestAllTypes", EMPTY_CONVERTER); + SelectField field = SelectField.create(TestAllTypes.SINGLE_INT64_FIELD_NUMBER, "single_int64"); - Optional nav = - raw.findByFieldNumber( - SelectField.create(TestAllTypes.SINGLE_INT64_FIELD_NUMBER, "single_int64")); + IllegalArgumentException thrown = + assertThrows(IllegalArgumentException.class, () -> raw.findByFieldNumber(field)); - assertThat(nav).hasValue(99L); + assertThat(thrown).hasCauseThat().hasMessageThat().contains("MESSAGE has unexpected wire type"); } @Test @@ -1438,8 +833,7 @@ public void selectByFieldNumber_mapEntrySpecDuplicateKeysAcrossEntries_lastEntry } @Test - public void selectByFieldNumber_mapEntrySpecUnexpectedFieldNumber_throwsIllegalStateException() - throws Exception { + public void selectByFieldNumber_mapEntrySpecUnknownFieldNumber_isSkipped() throws Exception { ByteString entry = encode( out -> { @@ -1458,10 +852,9 @@ public void selectByFieldNumber_mapEntrySpecUnexpectedFieldNumber_throwsIllegalS "map_string_string", SelectField.MapEntrySpec.create(9, 9)); - IllegalStateException thrown = - assertThrows(IllegalStateException.class, () -> raw.selectByFieldNumber(field)); + Object result = raw.selectByFieldNumber(field); - assertThat(thrown).hasMessageThat().contains("Unexpected field number in map entry: 3"); + assertThat(result).isEqualTo(ImmutableMap.of("k", "v")); } private static SelectField newMapInt64NestedTypeField() { @@ -1779,6 +1172,47 @@ public void selectByFieldNumber_fiveByteUnsignedInt32Varint_signExtendsCorrectly assertThat(result).isEqualTo(-42L); } + @Test + public void selectByFieldNumber_enumHighBits_truncatedToSigned32Bit() throws Exception { + ByteString wire = + encode(out -> out.writeUInt64(TestAllTypes.STANDALONE_ENUM_FIELD_NUMBER, 0x1FFFFFFFBL)); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + wire, "cel.expr.conformance.proto3.TestAllTypes", EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.STANDALONE_ENUM_FIELD_NUMBER, + "standalone_enum", + FieldLiteDescriptor.Type.ENUM.getNumber(), + 0L); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isEqualTo(-5L); + } + + @Test + public void selectByFieldNumber_invalidUtf8String_throwsIllegalArgumentException() + throws Exception { + byte[] invalidUtf8 = {(byte) 0xC0, (byte) 0xAF}; + ByteString wire = + encode(out -> out.writeByteArray(TestAllTypes.SINGLE_STRING_FIELD_NUMBER, invalidUtf8)); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + wire, "cel.expr.conformance.proto3.TestAllTypes", EMPTY_CONVERTER); + SelectField field = + SelectField.create( + TestAllTypes.SINGLE_STRING_FIELD_NUMBER, + "single_string", + FieldLiteDescriptor.Type.STRING.getNumber(), + ""); + + IllegalArgumentException thrown = + assertThrows(IllegalArgumentException.class, () -> raw.selectByFieldNumber(field)); + + assertThat(thrown).hasCauseThat().hasMessageThat().contains("invalid UTF-8"); + } + @Test public void selectByFieldNumber_unsetSubmessageWithProtoTypeName_returnsDefaultWithTypeName() { RawProtoMessageLiteValue raw =