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 @@ -46,6 +46,7 @@
import java.util.Map;
import java.util.Optional;
import java.util.TreeMap;
import org.jspecify.annotations.Nullable;

/**
* {@code ProtoLiteCelValueConverter} handles bidirectional conversion between native Java and
Expand All @@ -60,8 +61,8 @@
@Immutable
@Internal
public final class ProtoLiteCelValueConverter extends BaseProtoCelValueConverter {
static final String MAP_KEY_FIELD_NAME = "key";
static final String MAP_VALUE_FIELD_NAME = "value";
private static final String MAP_KEY_FIELD_NAME = "key";
private static final String MAP_VALUE_FIELD_NAME = "value";

private final CelLiteDescriptorPool descriptorPool;

Expand Down Expand Up @@ -134,29 +135,58 @@ private static Object readFixed64BitField(
}

private Object readLengthDelimitedField(
CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException {
CodedInputStream inputStream,
FieldLiteDescriptor fieldDescriptor,
@Nullable Object existingValue)
throws IOException {
FieldLiteDescriptor.Type fieldType = fieldDescriptor.getProtoFieldType();

switch (fieldType) {
case BYTES:
return inputStream.readBytes();
case MESSAGE:
String fieldProtoTypeName = fieldDescriptor.getFieldProtoTypeName();
MessageLiteDescriptor descriptor =
descriptorPool.findDescriptor(fieldProtoTypeName).orElse(null);
if (descriptor == null) {
return RawProtoMessageLiteValue.create(inputStream.readBytes(), fieldProtoTypeName, this);
}
MessageLite.Builder builder = descriptor.newMessageBuilder();
inputStream.readMessage(builder, ExtensionRegistryLite.getEmptyRegistry());
return builder.build();
return mergeOrReadMessageField(
inputStream.readBytes(), fieldDescriptor.getFieldProtoTypeName(), existingValue);
case STRING:
return inputStream.readStringRequireUtf8();
default:
throw new IllegalStateException("Unexpected field type: " + fieldType);
}
}

private Object readMessageField(ByteString bytes, String fieldProtoTypeName) {
return mergeOrReadMessageField(bytes, fieldProtoTypeName, /* existingValue= */ null);
}

private Object mergeOrReadMessageField(
ByteString bytes, String fieldProtoTypeName, @Nullable Object existingValue) {
MessageLiteDescriptor descriptor =
descriptorPool.findDescriptor(fieldProtoTypeName).orElse(null);
if (descriptor == null) {
if (existingValue instanceof RawProtoMessageLiteValue) {
bytes = ((RawProtoMessageLiteValue) existingValue).toByteString().concat(bytes);
}
return RawProtoMessageLiteValue.create(bytes, fieldProtoTypeName, this);
}
WellKnownProto wellKnownProto = WellKnownProto.getByTypeName(fieldProtoTypeName).orElse(null);
if (isStructLike(wellKnownProto)) {
if (existingValue instanceof ProtoMessageLiteValue) {
bytes = ((ProtoMessageLiteValue) existingValue).toByteString().concat(bytes);
}
return ProtoMessageLiteValue.create(bytes, fieldProtoTypeName, this);
}
if (existingValue instanceof MessageLite) {
return mergeMessageLite(((MessageLite) existingValue).toBuilder(), bytes, fieldProtoTypeName);
}
return parseMessageLite(bytes, descriptor);
}

// Unlike other WellKnownProtos (which unbox to CEL scalars/containers), FieldMask is
// represented as a standard struct message so field selection (e.g., mask.paths) works.
private static boolean isStructLike(@Nullable WellKnownProto wellKnownProto) {
return wellKnownProto == null || wellKnownProto == WellKnownProto.FIELD_MASK;
}

Object getDefaultCelValue(String protoTypeName, String fieldName) {
MessageLiteDescriptor messageDescriptor = descriptorPool.getDescriptorOrThrow(protoTypeName);
return getDefaultCelValue(messageDescriptor.getByFieldNameOrThrow(fieldName));
Expand All @@ -172,6 +202,28 @@ Optional<FieldLiteDescriptor> findFieldDescriptor(String protoTypeName, int fiel
.flatMap(desc -> desc.findByFieldNumber(fieldNumber));
}

MessageLite parseMessageLite(ByteString bytes, String protoTypeName) {
MessageLiteDescriptor descriptor = descriptorPool.getDescriptorOrThrow(protoTypeName);
return parseMessageLite(bytes, descriptor);
}

private static MessageLite parseMessageLite(ByteString bytes, MessageLiteDescriptor descriptor) {
if (bytes.isEmpty()) {
return descriptor.newMessageBuilder().getDefaultInstanceForType();
}
return mergeMessageLite(descriptor.newMessageBuilder(), bytes, descriptor.getProtoTypeName());
}

private static MessageLite mergeMessageLite(
MessageLite.Builder builder, ByteString bytes, String protoTypeName) {
try {
return builder.mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()).build();
} catch (IOException e) {
throw new IllegalArgumentException(
"Failed to decode proto message of type: " + protoTypeName, e);
}
}

Optional<Object> tryDecodeProtoMessage(ByteString bytes, String protoTypeName) {
return descriptorPool
.findDescriptor(protoTypeName)
Expand All @@ -180,14 +232,11 @@ Optional<Object> tryDecodeProtoMessage(ByteString bytes, String protoTypeName) {

private Object decodeProtoMessage(
ByteString bytes, String protoTypeName, MessageLiteDescriptor descriptor) {
try {
MessageLite.Builder builder =
descriptor.newMessageBuilder().mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry());
return toRuntimeValue(builder.build(), descriptor);
} catch (IOException e) {
throw new IllegalArgumentException(
"Failed to decode proto message of type: " + protoTypeName, e);
WellKnownProto wellKnownProto = WellKnownProto.getByTypeName(protoTypeName).orElse(null);
if (isStructLike(wellKnownProto)) {
return ProtoMessageLiteValue.create(bytes, protoTypeName, this);
}
return fromWellKnownProto(parseMessageLite(bytes, descriptor), checkNotNull(wellKnownProto));
}

@Override
Expand All @@ -209,12 +258,11 @@ public Object toRuntimeValue(Object value) {
private Object toRuntimeValue(MessageLite msg, MessageLiteDescriptor descriptor) {
WellKnownProto wellKnownProto =
WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null);

if (wellKnownProto == null || wellKnownProto == WellKnownProto.FIELD_MASK) {
if (isStructLike(wellKnownProto)) {
return ProtoMessageLiteValue.create(msg, descriptor.getProtoTypeName(), this);
}

return fromWellKnownProto(msg, wellKnownProto);
return fromWellKnownProto(msg, checkNotNull(wellKnownProto));
}

private Object getDefaultValue(FieldLiteDescriptor fieldDescriptor) {
Expand Down Expand Up @@ -260,12 +308,7 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) {
if (WellKnownProto.isWrapperType(fieldProtoTypeName)) {
return NullValue.NULL_VALUE;
}
MessageLiteDescriptor descriptor =
descriptorPool.findDescriptor(fieldProtoTypeName).orElse(null);
if (descriptor == null) {
return RawProtoMessageLiteValue.create(ByteString.EMPTY, fieldProtoTypeName, this);
}
return descriptor.newMessageBuilder().build();
return readMessageField(ByteString.EMPTY, fieldProtoTypeName);
}
throw new IllegalStateException("Unexpected java type: " + type);
}
Expand All @@ -286,7 +329,7 @@ private Map.Entry<Object, Object> readSingleMapEntry(
CodedInputStream inputStream, FieldLiteDescriptor fieldDescriptor) throws IOException {
String entryTypeName = fieldDescriptor.getFieldProtoTypeName();
ImmutableMap<String, Object> singleMapEntry =
readAllFields(inputStream.readByteArray(), entryTypeName).values();
readAllFields(inputStream.readBytes(), entryTypeName).values();
Object key = singleMapEntry.get(MAP_KEY_FIELD_NAME);
if (key == null) {
key = getDefaultCelValue(entryTypeName, MAP_KEY_FIELD_NAME);
Expand All @@ -299,25 +342,32 @@ private Map.Entry<Object, Object> readSingleMapEntry(
return new AbstractMap.SimpleEntry<>(key, value);
}

MessageFields readAllFields(byte[] bytes, String protoTypeName) throws IOException {
MessageFields readAllFields(ByteString bytes, String protoTypeName) throws IOException {
MessageLiteDescriptor messageDescriptor = descriptorPool.getDescriptorOrThrow(protoTypeName);
CodedInputStream inputStream = CodedInputStream.newInstance(bytes);
if (bytes.isEmpty()) {
return MessageFields.EMPTY;
}
return readAllFields(bytes.newCodedInput(), messageDescriptor);
}

Multimap<Integer, Object> unknownFields =
Multimaps.newMultimap(new TreeMap<>(), ArrayList::new);
ImmutableMap.Builder<String, Object> fieldValues = ImmutableMap.builder();
Map<Integer, List<Object>> repeatedFieldValues = new LinkedHashMap<>();
Map<Integer, Map<Object, Object>> mapFieldValues = new LinkedHashMap<>();
private MessageFields readAllFields(
CodedInputStream inputStream, MessageLiteDescriptor messageDescriptor) throws IOException {
Multimap<Integer, Object> unknownFields = null;
Map<String, Object> fieldValues = new LinkedHashMap<>();
for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) {
int tagWireType = WireFormat.getTagWireType(tag);
int fieldNumber = WireFormat.getTagFieldNumber(tag);
FieldLiteDescriptor fieldDescriptor =
messageDescriptor.findByFieldNumber(fieldNumber).orElse(null);
if (fieldDescriptor == null) {
if (unknownFields == null) {
unknownFields = Multimaps.newMultimap(new TreeMap<>(), ArrayList::new);
}
unknownFields.put(fieldNumber, readUnknownField(tagWireType, inputStream));
continue;
}

String fieldName = fieldDescriptor.getFieldName();
Object payload;
switch (tagWireType) {
case WireFormat.WIRETYPE_VARINT:
Expand Down Expand Up @@ -347,18 +397,24 @@ MessageFields readAllFields(byte[] bytes, String protoTypeName) throws IOExcepti
+ protoFieldType);
}

payload = readLengthDelimitedField(inputStream, fieldDescriptor);
payload =
readLengthDelimitedField(
inputStream, fieldDescriptor, /* existingValue= */ null);
}
break;
case MAP:
// Safe because MAP fields only ever store a LinkedHashMap in fieldValues.
@SuppressWarnings("unchecked")
Map<Object, Object> fieldMap =
mapFieldValues.computeIfAbsent(fieldNumber, (unused) -> new LinkedHashMap<>());
(Map<Object, Object>)
fieldValues.computeIfAbsent(fieldName, (unused) -> new LinkedHashMap<>());
Map.Entry<Object, Object> mapEntry = readSingleMapEntry(inputStream, fieldDescriptor);
fieldMap.put(mapEntry.getKey(), mapEntry.getValue());
payload = fieldMap;
break;
continue;
default:
payload = readLengthDelimitedField(inputStream, fieldDescriptor);
payload =
readLengthDelimitedField(
inputStream, fieldDescriptor, fieldValues.get(fieldName));
break;
}
break;
Expand All @@ -371,30 +427,27 @@ MessageFields readAllFields(byte[] bytes, String protoTypeName) throws IOExcepti
}

if (fieldDescriptor.getEncodingType().equals(EncodingType.LIST)) {
String fieldName = fieldDescriptor.getFieldName();
List<Object> repeatedValues =
repeatedFieldValues.computeIfAbsent(fieldNumber, (unused) -> new ArrayList<>());

if (payload instanceof Collection) {
repeatedValues.addAll((Collection<?>) payload);
Collection<?> elements = (Collection<?>) payload;
if (!elements.isEmpty()) {
getOrCreateRepeatedList(fieldValues, fieldName).addAll(elements);
}
} else {
repeatedValues.add(payload);
}
if (!repeatedValues.isEmpty()) {
fieldValues.put(fieldName, repeatedValues);
getOrCreateRepeatedList(fieldValues, fieldName).add(payload);
}
} else {
fieldValues.put(fieldDescriptor.getFieldName(), payload);
fieldValues.put(fieldName, payload);
}
}

// Protobuf encoding follows a "last one wins" semantics. This means for duplicated fields,
// we accept the last value encountered.
return MessageFields.create(fieldValues.buildKeepingLast(), unknownFields);
return MessageFields.create(ImmutableMap.copyOf(fieldValues), unknownFields);
}

MessageFields readMessageFields(MessageLite msg, String protoTypeName) throws IOException {
return readAllFields(msg.toByteArray(), protoTypeName);
// Safe because LIST fields only ever store an ArrayList in fieldValues.
@SuppressWarnings("unchecked")
private static List<Object> getOrCreateRepeatedList(
Map<String, Object> fieldValues, String fieldName) {
return (List<Object>) fieldValues.computeIfAbsent(fieldName, (unused) -> new ArrayList<>());
}

static Object readUnknownField(int tagWireType, CodedInputStream inputStream) throws IOException {
Expand All @@ -421,19 +474,26 @@ static Object readUnknownField(int tagWireType, CodedInputStream inputStream) th
@Immutable
@SuppressWarnings("Immutable") // Safe immutable fields
abstract static class MessageFields {
static final MessageFields EMPTY =
new AutoValue_ProtoLiteCelValueConverter_MessageFields(
ImmutableMap.of(), ImmutableListMultimap.of());

abstract ImmutableMap<String, Object> values();

abstract ImmutableListMultimap<Integer, Object> unknowns();

static MessageFields create(
ImmutableMap<String, Object> fieldValues, Multimap<Integer, Object> unknownFields) {
private static MessageFields create(
ImmutableMap<String, Object> fieldValues,
@Nullable Multimap<Integer, Object> unknownFields) {
return new AutoValue_ProtoLiteCelValueConverter_MessageFields(
fieldValues, ImmutableListMultimap.copyOf(unknownFields));
fieldValues,
unknownFields == null
? ImmutableListMultimap.of()
: ImmutableListMultimap.copyOf(unknownFields));
}
}

private ProtoLiteCelValueConverter(CelLiteDescriptorPool celLiteDescriptorPool) {
this.descriptorPool = celLiteDescriptorPool;
this.descriptorPool = checkNotNull(celLiteDescriptorPool);
}
}
Loading
Loading