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.

134 changes: 112 additions & 22 deletions common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java
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,41 +40,72 @@
* <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>
* <li><b>Field renames:</b> If a protobuf field is renamed in schema after an AST was compiled,
* resolving by {@link SelectField#fieldNumber()} maps the number to the runtime descriptor's
* current field name, preventing {@code CelAttributeNotFoundException}.
* <li><b>Version skew / unknown fields:</b> When evaluating payloads serialized by a newer binary
* containing fields absent from the local {@code CelLiteDescriptor}, the unknown wire bytes
* are preserved in {@link #unknownFields()} and decoded on demand using the compile-time wire
* type and default metadata in {@link SelectField}.
* containing fields absent from the local {@code CelLiteDescriptor}, the unknown fields are
* decoded on demand directly from the message's wire bytes using the compile-time wire type
* and default metadata in {@link SelectField}.
* </ul>
*/
@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());
}

@Memoized
ByteString serializedRawValue() {
return checkNotNull(rawValue()).toByteString();
}

ByteString toByteString() {
ByteString bytes = wireBytes();
return bytes != null ? bytes : serializedRawValue();
}

@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 +115,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,16 +147,15 @@ 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
public Object selectByFieldNumber(SelectField field) {
FieldLiteDescriptor fd = findFieldDescriptor(field);
if (fd != null) {
Object known = fieldValues().get(fd.getFieldName());
Object known = readField(fd);
if (known != null) {
return protoLiteCelValueConverter().toRuntimeValue(known);
}
Expand All @@ -112,28 +165,48 @@ public Object selectByFieldNumber(SelectField field) {
return protoLiteCelValueConverter().getDefaultCelValue(fd);
}
return RawProtoMessageLiteValue.selectWireOrDefault(
field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter());
field,
RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()),
protoLiteCelValueConverter());
}

@Override
public boolean hasFieldByNumber(SelectField field) {
FieldLiteDescriptor fd = findFieldDescriptor(field);
if (fd != null) {
return fieldValues().containsKey(fd.getFieldName());
return hasField(fd);
}
return RawProtoMessageLiteValue.isPresentInWire(
field, unknownFields().get(field.fieldNumber()));
return RawProtoMessageLiteValue.isPresentInWire(toByteString(), field);
}

@Override
public Optional<Object> findByFieldNumber(SelectField field) {
FieldLiteDescriptor fd = findFieldDescriptor(field);
if (fd != null) {
return Optional.ofNullable(fieldValues().get(fd.getFieldName()))
.map(value -> protoLiteCelValueConverter().toRuntimeValue(value));
return Optional.ofNullable(readField(fd)).map(protoLiteCelValueConverter()::toRuntimeValue);
}
return RawProtoMessageLiteValue.navigateWire(
field, unknownFields().get(field.fieldNumber()), protoLiteCelValueConverter());
field,
RawProtoMessageLiteValue.readWireField(toByteString(), field.fieldNumber()),
protoLiteCelValueConverter());
}

private @Nullable Object readField(FieldLiteDescriptor fd) {
try {
return protoLiteCelValueConverter().readSingleField(toByteString(), fd);
} catch (IOException e) {
throw new IllegalArgumentException(
"Failed to decode proto message of type: " + celType().name(), e);
}
}

private boolean hasField(FieldLiteDescriptor fd) {
try {
return protoLiteCelValueConverter().hasSingleField(toByteString(), fd);
} catch (IOException e) {
throw new IllegalArgumentException(
"Failed to decode proto message of type: " + celType().name(), e);
}
}

private @Nullable FieldLiteDescriptor findFieldDescriptor(SelectField field) {
Expand All @@ -142,13 +215,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