Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -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;

Expand Down Expand Up @@ -194,21 +198,6 @@ private static MessageLite mergeMessageLite(
}
}

Optional<Object> 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);
Expand Down Expand Up @@ -284,12 +273,10 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) {
}

private Map.Entry<Object, Object> 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;
Expand Down Expand Up @@ -337,14 +324,9 @@ private Map.Entry<Object, Object> 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);
Expand All @@ -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<Object> 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<Object, Object> 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.
Expand Down Expand Up @@ -520,18 +622,38 @@ private Object readSingularField(
return repeatedValues;
}

private Map<Object, Object> 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<Object, Object> readMapField(
int tagWireType,
CodedInputStream inputStream,
FieldLiteDescriptor fieldDescriptor,
FieldLiteDescriptor keyDescriptor,
FieldLiteDescriptor valueDescriptor,
@Nullable Object existingValue)
throws IOException {
checkWireType(tagWireType, fieldDescriptor);
Map<Object, Object> mapValues =
existingValue != null ? (Map<Object, Object>) existingValue : new LinkedHashMap<>();
Map.Entry<Object, Object> mapEntry = readSingleMapEntry(inputStream, fieldDescriptor);
Map.Entry<Object, Object> mapEntry =
readSingleMapEntry(inputStream, keyDescriptor, valueDescriptor);
mapValues.put(mapEntry.getKey(), mapEntry.getValue());
return mapValues;
}
Expand Down Expand Up @@ -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) {
Expand All @@ -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}.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -175,10 +182,12 @@ public Optional<Object> 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) {
Expand Down
Loading
Loading