Skip to content
Open
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
4 changes: 4 additions & 0 deletions common/src/main/java/dev/cel/common/values/BUILD.bazel
Original file line number Diff line number Diff line change
Expand Up @@ -280,7 +280,9 @@ java_library(
],
deps = [
":base_proto_cel_value_converter",
":optimized_selectable",
":preadapted_list",
":select_field",
":values",
"//:auto_value",
"//common:options",
Expand Down Expand Up @@ -320,6 +322,7 @@ java_library(
"ProtoLiteCelValueConverter.java",
"ProtoMessageLiteValue.java",
"RawProtoMessageLiteValue.java",
"WireMessageLite.java",
],
tags = [
],
Expand Down Expand Up @@ -350,6 +353,7 @@ cel_android_library(
"ProtoLiteCelValueConverter.java",
"ProtoMessageLiteValue.java",
"RawProtoMessageLiteValue.java",
"WireMessageLite.java",
],
tags = [
],
Expand Down

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -21,12 +21,14 @@
import com.google.common.collect.ImmutableListMultimap;
import com.google.common.collect.ImmutableMap;
import com.google.errorprone.annotations.Immutable;
import com.google.protobuf.ByteString;
import com.google.protobuf.MessageLite;
import dev.cel.common.types.CelType;
import dev.cel.common.types.StructTypeReference;
import dev.cel.common.values.ProtoLiteCelValueConverter.MessageFields;
import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor;
import java.io.IOException;
import java.util.Objects;
import java.util.Optional;
import org.jspecify.annotations.Nullable;

Expand All @@ -38,6 +40,11 @@
* <p>If the codebase has access to full protobuf messages with descriptors, use {@code
* ProtoMessageValue} instead.
*
* <p>An instance is backed by either a materialized {@link #rawValue()} or unparsed {@link
* #wireBytes()} (exactly one is non-null). Wire-backed instances decode selected fields directly
* from {@link #wireBytes()} and lazily parse the full {@link MessageLite} only if {@link #value()}
* is invoked.
*
* <p>Implements {@link OptimizedSelectable} so that select chains can address fields by number:
*
* <ul>
Expand All @@ -52,27 +59,48 @@
*/
@AutoValue
@Immutable
public abstract class ProtoMessageLiteValue extends StructValue<String, MessageLite>
abstract class ProtoMessageLiteValue extends StructValue<String, MessageLite>
implements OptimizedSelectable {

@Override
public abstract MessageLite value();
// Populated when wrapping an already-materialized root MessageLite (e.g., from activation).
abstract @Nullable MessageLite rawValue();

// Populated when slicing a nested submessage from parent wire bytes to avoid deserializing and
// re-serializing intermediate hops; lazily parsed into a MessageLite only if value() is called.
abstract @Nullable ByteString wireBytes();

@Override
public abstract CelType celType();

abstract ProtoLiteCelValueConverter protoLiteCelValueConverter();

@Memoized
@Override
public MessageLite value() {
MessageLite msg = rawValue();
if (msg != null) {
return msg;
}
return protoLiteCelValueConverter()
.parseMessageLite(checkNotNull(wireBytes()), celType().name());
}

ByteString toByteString() {
ByteString bytes = wireBytes();
return bytes != null ? bytes : checkNotNull(rawValue()).toByteString();
}

@Memoized
MessageFields messageFields() {
try {
return protoLiteCelValueConverter().readMessageFields(value(), celType().name());
return protoLiteCelValueConverter().readAllFields(toByteString(), celType().name());
} catch (IOException e) {
throw new IllegalStateException("Unable to read message fields for " + celType().name(), e);
throw new IllegalArgumentException(
"Failed to decode proto message of type: " + celType().name(), e);
}
}

ImmutableMap<String, Object> fieldValues() {
private ImmutableMap<String, Object> fieldValues() {
return messageFields().values();
}

Expand All @@ -82,9 +110,30 @@ ImmutableListMultimap<Integer, Object> unknownFields() {

@Override
public boolean isZeroValue() {
ByteString bytes = wireBytes();
if (bytes != null && bytes.isEmpty()) {
return true;
}
return value().getDefaultInstanceForType().equals(value());
}

@Override
public final boolean equals(Object other) {
if (other == this) {
return true;
}
if (!(other instanceof ProtoMessageLiteValue)) {
return false;
}
ProtoMessageLiteValue that = (ProtoMessageLiteValue) other;
return this.celType().equals(that.celType()) && this.value().equals(that.value());
}

@Override
public final int hashCode() {
return Objects.hash(value(), celType());
}

@Override
public Object select(String field) {
return find(field)
Expand All @@ -93,9 +142,8 @@ public Object select(String field) {

@Override
public Optional<Object> find(String field) {
Object fieldValue = fieldValues().get(field);
return Optional.ofNullable(fieldValue)
.map(value -> protoLiteCelValueConverter().toRuntimeValue(fieldValue));
return Optional.ofNullable(fieldValues().get(field))
.map(protoLiteCelValueConverter()::toRuntimeValue);
}

@Override
Expand Down Expand Up @@ -130,7 +178,7 @@ public Optional<Object> findByFieldNumber(SelectField field) {
FieldLiteDescriptor fd = findFieldDescriptor(field);
if (fd != null) {
return Optional.ofNullable(fieldValues().get(fd.getFieldName()))
.map(value -> protoLiteCelValueConverter().toRuntimeValue(value));
.map(protoLiteCelValueConverter()::toRuntimeValue);
}
return RawProtoMessageLiteValue.navigateWire(
field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter());
Expand All @@ -142,13 +190,30 @@ public Optional<Object> findByFieldNumber(SelectField field) {
.orElse(null);
}

public static ProtoMessageLiteValue create(
static ProtoMessageLiteValue create(
MessageLite value, String typeName, ProtoLiteCelValueConverter protoLiteCelValueConverter) {
checkNotNull(value);
checkNotNull(typeName);
checkNotNull(protoLiteCelValueConverter);
return new AutoValue_ProtoMessageLiteValue(
value, StructTypeReference.create(typeName), protoLiteCelValueConverter);
value,
/* wireBytes= */ null,
StructTypeReference.create(typeName),
protoLiteCelValueConverter);
}

static ProtoMessageLiteValue create(
ByteString wireBytes,
String typeName,
ProtoLiteCelValueConverter protoLiteCelValueConverter) {
checkNotNull(wireBytes);
checkNotNull(typeName);
checkNotNull(protoLiteCelValueConverter);
return new AutoValue_ProtoMessageLiteValue(
/* rawValue= */ null,
wireBytes,
StructTypeReference.create(typeName),
protoLiteCelValueConverter);
}

ProtoMessageLiteValue() {}
Expand Down
76 changes: 60 additions & 16 deletions common/src/main/java/dev/cel/common/values/ProtoMessageValue.java
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,9 @@

package dev.cel.common.values;

import static com.google.common.base.Preconditions.checkNotNull;

import com.google.auto.value.AutoValue;
import com.google.common.base.Preconditions;
import com.google.errorprone.annotations.Immutable;
import com.google.protobuf.Descriptors.Descriptor;
import com.google.protobuf.Descriptors.FieldDescriptor;
Expand All @@ -28,7 +29,8 @@
/** ProtoMessageValue is a struct value with protobuf support. */
@AutoValue
@Immutable
public abstract class ProtoMessageValue extends StructValue<String, Message> {
public abstract class ProtoMessageValue extends StructValue<String, Message>
implements OptimizedSelectable {

@Override
public abstract Message value();
Expand Down Expand Up @@ -60,28 +62,38 @@ public Optional<Object> find(String field) {
FieldDescriptor fieldDescriptor =
findField(celDescriptorPool(), value().getDescriptorForType(), field);

// Selecting a field on a protobuf message yields a default value even if the field is not
// declared. Therefore, we must exhaustively test whether they are actually declared.
if (fieldDescriptor.isRepeated()) {
if (value().getRepeatedFieldCount(fieldDescriptor) == 0) {
return Optional.empty();
}
} else if (!value().hasField(fieldDescriptor)) {
return Optional.empty();
}
return findFieldValue(fieldDescriptor);
}

return Optional.of(
protoCelValueConverter().fromProtoMessageFieldToCelValue(value(), fieldDescriptor));
@Override
public Object selectByFieldNumber(SelectField field) {
FieldDescriptor fieldDescriptor = findFieldByNumber(value().getDescriptorForType(), field);

return protoCelValueConverter().fromProtoMessageFieldToCelValue(value(), fieldDescriptor);
}

@Override
public boolean hasFieldByNumber(SelectField field) {
FieldDescriptor fieldDescriptor = findFieldByNumber(value().getDescriptorForType(), field);

return isFieldPresent(fieldDescriptor);
}

@Override
public Optional<Object> findByFieldNumber(SelectField field) {
FieldDescriptor fieldDescriptor = findFieldByNumber(value().getDescriptorForType(), field);

return findFieldValue(fieldDescriptor);
}

public static ProtoMessageValue create(
Message value,
CelDescriptorPool celDescriptorPool,
ProtoCelValueConverter protoCelValueConverter,
boolean enableJsonFieldNames) {
Preconditions.checkNotNull(value);
Preconditions.checkNotNull(celDescriptorPool);
Preconditions.checkNotNull(protoCelValueConverter);
checkNotNull(value);
checkNotNull(celDescriptorPool);
checkNotNull(protoCelValueConverter);
return new AutoValue_ProtoMessageValue(
value,
StructTypeReference.create(value.getDescriptorForType().getFullName()),
Expand All @@ -90,6 +102,36 @@ public static ProtoMessageValue create(
enableJsonFieldNames);
}

private Optional<Object> findFieldValue(FieldDescriptor fieldDescriptor) {
if (!isFieldPresent(fieldDescriptor)) {
return Optional.empty();
}

return Optional.of(
protoCelValueConverter().fromProtoMessageFieldToCelValue(value(), fieldDescriptor));
}

private boolean isFieldPresent(FieldDescriptor fieldDescriptor) {
// Selecting a field on a protobuf message yields a default value even if the field is not
// declared. Therefore, we must exhaustively test whether they are actually declared.
if (fieldDescriptor.isRepeated()) {
return value().getRepeatedFieldCount(fieldDescriptor) > 0;
}
return value().hasField(fieldDescriptor);
}

private static FieldDescriptor findFieldByNumber(Descriptor descriptor, SelectField field) {
FieldDescriptor fieldDescriptor = descriptor.findFieldByNumber(field.fieldNumber());
if (fieldDescriptor != null) {
return fieldDescriptor;
}

throw new IllegalArgumentException(
String.format(
"field '%s' (number %d) is not declared in message '%s'",
field.fieldName(), field.fieldNumber(), descriptor.getFullName()));
}

private FieldDescriptor findField(
CelDescriptorPool celDescriptorPool, Descriptor descriptor, String fieldName) {
if (enableJsonFieldNames()) {
Expand All @@ -114,4 +156,6 @@ private FieldDescriptor findField(
"field '%s' is not declared in message '%s'",
fieldName, descriptor.getFullName())));
}

ProtoMessageValue() {}
}
Loading
Loading