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
2 changes: 2 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 @@ -320,6 +320,7 @@ java_library(
"ProtoLiteCelValueConverter.java",
"ProtoMessageLiteValue.java",
"RawProtoMessageLiteValue.java",
"WireMessageLite.java",
],
tags = [
],
Expand Down Expand Up @@ -350,6 +351,7 @@ cel_android_library(
"ProtoLiteCelValueConverter.java",
"ProtoMessageLiteValue.java",
"RawProtoMessageLiteValue.java",
"WireMessageLite.java",
],
tags = [
],
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,6 @@
import com.google.protobuf.CodedInputStream;
import com.google.protobuf.ExtensionRegistryLite;
import com.google.protobuf.MessageLite;
import com.google.protobuf.MessageLiteOrBuilder;
import com.google.protobuf.WireFormat;
import dev.cel.common.annotations.Internal;
import dev.cel.common.internal.CelLiteDescriptorPool;
Expand All @@ -45,7 +44,6 @@
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.NoSuchElementException;
import java.util.Optional;
import java.util.TreeMap;

Expand Down Expand Up @@ -110,6 +108,8 @@ private static Object readFixed32BitField(
case FLOAT:
return inputStream.readFloat();
case FIXED32:
return UnsignedLong.fromLongBits(
Integer.toUnsignedLong(inputStream.readRawLittleEndian32()));
case SFIXED32:
return inputStream.readRawLittleEndian32();
default:
Expand All @@ -124,6 +124,7 @@ private static Object readFixed64BitField(
case DOUBLE:
return inputStream.readDouble();
case FIXED64:
return UnsignedLong.fromLongBits(inputStream.readRawLittleEndian64());
case SFIXED64:
return inputStream.readRawLittleEndian64();
default:
Expand Down Expand Up @@ -171,31 +172,21 @@ Optional<FieldLiteDescriptor> findFieldDescriptor(String protoTypeName, int fiel
.flatMap(desc -> desc.findByFieldNumber(fieldNumber));
}

Optional<Object> tryDecodeWellKnownProto(ByteString bytes, String protoTypeName) {
Optional<WellKnownProto> wellKnownProto = WellKnownProto.getByTypeName(protoTypeName);
if (!wellKnownProto.isPresent()) {
return Optional.empty();
}

Optional<Object> tryDecodeProtoMessage(ByteString bytes, String protoTypeName) {
return descriptorPool
.findDescriptor(protoTypeName)
.map(
descriptor ->
decodeWellKnownProto(bytes, protoTypeName, descriptor, wellKnownProto.get()));
.map(descriptor -> decodeProtoMessage(bytes, protoTypeName, descriptor));
}

private Object decodeWellKnownProto(
ByteString bytes,
String protoTypeName,
MessageLiteDescriptor descriptor,
WellKnownProto wellKnownProto) {
private Object decodeProtoMessage(
ByteString bytes, String protoTypeName, MessageLiteDescriptor descriptor) {
try {
MessageLite.Builder builder = descriptor.newMessageBuilder();
builder.mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry());
return fromWellKnownProto(builder.build(), wellKnownProto);
MessageLite.Builder builder =
descriptor.newMessageBuilder().mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry());
return toRuntimeValue(builder.build(), descriptor);
} catch (IOException e) {
throw new IllegalArgumentException(
"Failed to decode well-known proto of type: " + protoTypeName, e);
"Failed to decode proto message of type: " + protoTypeName, e);
}
}

Expand All @@ -209,35 +200,21 @@ public Object toRuntimeValue(Object value) {
if (descriptor == null) {
return RawProtoMessageLiteValue.create(msg.toByteString(), this);
}
WellKnownProto wellKnownProto =
WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null);

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

return fromWellKnownProto(msg, wellKnownProto);
return toRuntimeValue(msg, descriptor);
}

return super.toRuntimeValue(value);
}

@Override
protected Object fromWellKnownProto(MessageLiteOrBuilder msg, WellKnownProto wellKnownProto) {
if (wellKnownProto == WellKnownProto.FIELD_MASK) {
MessageLite message = (MessageLite) msg;
MessageLiteDescriptor descriptor =
descriptorPool
.findDescriptor(message)
.orElseThrow(
() ->
new NoSuchElementException(
"Could not find a descriptor for message of type: "
+ message.getClass().getName()));
return ProtoMessageLiteValue.create(message, descriptor.getProtoTypeName(), this);
private Object toRuntimeValue(MessageLite msg, MessageLiteDescriptor descriptor) {
WellKnownProto wellKnownProto =
WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null);

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

return super.fromWellKnownProto(msg, wellKnownProto);
return fromWellKnownProto(msg, wellKnownProto);
}

private Object getDefaultValue(FieldLiteDescriptor fieldDescriptor) {
Expand All @@ -257,11 +234,13 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) {
JavaType type = fieldDescriptor.getJavaType();
switch (type) {
case INT:
return fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT32)
return (fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT32)
|| fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.FIXED32))
? UnsignedLong.ZERO
: Defaults.defaultValue(long.class);
case LONG:
return fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT64)
return (fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT64)
|| fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.FIXED64))
? UnsignedLong.ZERO
: Defaults.defaultValue(long.class);
case ENUM:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@
*/
@AutoValue
@Immutable
public abstract class ProtoMessageLiteValue extends StructValue<String, MessageLite>
abstract class ProtoMessageLiteValue extends StructValue<String, MessageLite>
implements OptimizedSelectable {

@Override
Expand Down Expand Up @@ -142,7 +142,7 @@ 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);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,6 @@
import com.google.protobuf.ByteString;
import com.google.protobuf.CodedInputStream;
import com.google.protobuf.WireFormat;
import dev.cel.common.annotations.Internal;
import dev.cel.common.exceptions.CelAttributeNotFoundException;
import dev.cel.common.types.CelType;
import dev.cel.common.types.StructTypeReference;
Expand All @@ -39,6 +38,7 @@
import java.util.AbstractMap;
import java.util.ArrayList;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Optional;
import java.util.TreeMap;
Expand All @@ -49,38 +49,63 @@
* client-server version skew issues where newer fields or submessages lack generated classes and
* descriptors in the evaluation environment.
*
* <p>Rather than requiring compiled {@link MessageLite} classes or runtime schema descriptors, this
* <p>Rather than requiring compiled {@code MessageLite} classes or runtime schema descriptors, this
* value encapsulates the raw wire-format {@link ByteString} payload and performs classless,
* reflection-free field traversal directly over wire tags via {@link CodedInputStream}.
*/
@AutoValue
@AutoValue.CopyAnnotations
@Immutable
@SuppressWarnings("Immutable") // Immutable wire fields
@Internal
public abstract class RawProtoMessageLiteValue extends StructValue<String, RawProtoMessageLiteValue>
implements OptimizedSelectable {
abstract class RawProtoMessageLiteValue extends StructValue<String, WireMessageLite>
implements OptimizedSelectable, WireMessageLite {

private static final String UNKNOWN_MESSAGE_TYPE_NAME = "cel.@unknownMessage";
private static final int MAP_KEY_FIELD_NUMBER = 1;
private static final int MAP_VALUE_FIELD_NUMBER = 2;

abstract ByteString rawWireBytes();
@Override
public abstract ByteString toByteString();

@Override
public abstract CelType celType();

abstract ProtoLiteCelValueConverter protoLiteCelValueConverter();

@Override
public RawProtoMessageLiteValue value() {
public String protoTypeName() {
return celType().name();
}

@Override
public WireMessageLite value() {
return this;
}

@Override
public final boolean equals(Object other) {
// TODO: Support message equality
throw new UnsupportedOperationException("Message equality is not supported");
}

@Override
public final int hashCode() {
throw new UnsupportedOperationException("Message equality is not supported");
}

@Override
public final String toString() {
return String.format(
Locale.US,
"WireMessageLite{protoTypeName=%s, size=%d}",
protoTypeName(),
toByteString().size());
}

@Memoized
ImmutableListMultimap<Integer, Object> unknownFields() {
try {
CodedInputStream inputStream = rawWireBytes().newCodedInput();
CodedInputStream inputStream = toByteString().newCodedInput();
Multimap<Integer, Object> fields = Multimaps.newMultimap(new TreeMap<>(), ArrayList::new);
for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) {
int tagWireType = WireFormat.getTagWireType(tag);
Expand All @@ -96,7 +121,7 @@ ImmutableListMultimap<Integer, Object> unknownFields() {

@Override
public boolean isZeroValue() {
return rawWireBytes().isEmpty();
return toByteString().isEmpty();
}

/**
Expand Down Expand Up @@ -169,27 +194,15 @@ private static Object decodeWireField(
}

boolean isRepeated = field.defaultValue() instanceof List;
String protoTypeName = resolveProtoTypeName(field);

return decodeWireEntries(unknowns, typeCode, protoTypeName, isRepeated, converter);
}

/**
* Resolves the protobuf message type name for a field from the optimizer metadata in {@link
* SelectField}, or {@link #UNKNOWN_MESSAGE_TYPE_NAME} if unspecified.
*/
private static String resolveProtoTypeName(SelectField field) {
if (!field.protoTypeName().isEmpty()) {
return field.protoTypeName();
}
return UNKNOWN_MESSAGE_TYPE_NAME;
return decodeWireEntries(unknowns, typeCode, field.protoTypeName(), isRepeated, converter);
}

private static Object resolveDefault(SelectField field, ProtoLiteCelValueConverter converter) {
if (field.defaultValue() != null) {
return field.defaultValue();
}
return create(ByteString.EMPTY, resolveProtoTypeName(field), converter);
return decodeMessageValue(ByteString.EMPTY, field.protoTypeName(), converter);
}

/**
Expand Down Expand Up @@ -419,9 +432,7 @@ static Object decodeWireValue(
throw new UnsupportedOperationException("Groups are not supported");
case MESSAGE:
ByteString msgBytes = requireType(raw, ByteString.class, fieldType);
return converter
.tryDecodeWellKnownProto(msgBytes, protoTypeName)
.orElseGet(() -> create(msgBytes, protoTypeName, converter));
return decodeMessageValue(msgBytes, protoTypeName, converter);
case BYTES:
return CelByteString.of(requireType(raw, ByteString.class, fieldType).toByteArray());
case UINT32:
Expand All @@ -437,6 +448,13 @@ static Object decodeWireValue(
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> T requireType(
Object raw, Class<T> expectedType, WireFormat.FieldType fieldType) {
if (!expectedType.isInstance(raw)) {
Expand Down Expand Up @@ -509,12 +527,12 @@ private static ImmutableList<Object> decodePacked(
}
}

public static RawProtoMessageLiteValue create(
static RawProtoMessageLiteValue create(
ByteString rawWireBytes, ProtoLiteCelValueConverter protoLiteCelValueConverter) {
return create(rawWireBytes, "", protoLiteCelValueConverter);
}

public static RawProtoMessageLiteValue create(
static RawProtoMessageLiteValue create(
ByteString rawWireBytes,
String protoTypeName,
ProtoLiteCelValueConverter protoLiteCelValueConverter) {
Expand Down
47 changes: 47 additions & 0 deletions common/src/main/java/dev/cel/common/values/WireMessageLite.java
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
// Copyright 2026 Google LLC
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// https://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

package dev.cel.common.values;

import com.google.errorprone.annotations.Immutable;
import com.google.protobuf.ByteString;
import dev.cel.common.annotations.Beta;

/**
* Represents a protobuf message evaluation result in {@code CelLiteRuntime} when no {@code
* CelLiteDescriptor} is registered for the message type.
*
* <p>When a message-typed expression is evaluated in {@code CelLiteRuntime}:
*
* <ul>
* <li>If a {@code CelLiteDescriptor} is registered for the message type, evaluation produces a
* {@code MessageLite} instance.
* <li>Otherwise, evaluation produces a {@code WireMessageLite} carrying the message's protobuf
* type name and wire-encoded payload.
* </ul>
*/
@Immutable
@Beta
public interface WireMessageLite {

/**
* Returns the fully-qualified protobuf message type name (e.g. {@code
* "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"}), or {@code "cel.@unknownMessage"} if
* the message type name is not known at runtime.
*/
String protoTypeName();

/** Serializes the message to a {@link ByteString} in protobuf wire format. */
ByteString toByteString();
}
Loading
Loading